Skip to content

vllm_omni.diffusion.models.magi2.parallel

Distributed primitives for the native MAGI-2 preview transformer.

This file is adapted from SandAI's Apache-2.0 MAGI-2 preview context- and expert-parallel primitives. It has been modified to reuse vLLM-Omni's already-initialized sequence-parallel and expert-parallel process groups.

MAGI-2 supports two equivalent layouts for its two parallel views:

  • Ulysses context parallelism shards tokens and gathers attention heads.
  • Multi-head expert parallelism shards the MoE head axis, not experts.

With SP-only deployment both views use the SP group. With TP enabled, Ulysses keeps using SP while native column/row tensor parallelism and the MoE heads use TP. This gives TP4 and TP2SP2 independent weight/head and token axes while preserving the released equations and checkpoint hierarchy.

These helpers intentionally do not create or own process groups.

Magi2ParallelGroup dataclass

The small process-group surface needed by MAGI-2 kernels.

group instance-attribute

group: ProcessGroup | None

rank instance-attribute

rank: int

replicated_sequence class-attribute instance-attribute

replicated_sequence: bool = False

world_size instance-attribute

world_size: int

Magi2SequenceDispatcher

Request-scoped dispatcher that enforces one consistent token split.

group instance-attribute

group = group or get_magi2_ulysses_group()

split_sizes instance-attribute

split_sizes: list[int] | None = None

dispatch

dispatch(tensor: Tensor) -> Tensor

undispatch

undispatch(tensor: Tensor) -> Tensor

balanced_split_sizes

balanced_split_sizes(
    length: int, world_size: int
) -> list[int]

Split length contiguously, putting one extra token on early ranks.

ep_dispatch

ep_dispatch(
    tensor: Tensor,
    group: Magi2ParallelGroup | None = None,
    sequence_split_sizes: list[int] | None = None,
) -> Tensor

Dispatch [S,H,D] so each rank evaluates a contiguous head shard.

ep_undispatch

ep_undispatch(
    tensor: Tensor,
    group: Magi2ParallelGroup | None = None,
    sequence_split_sizes: list[int] | None = None,
) -> Tensor

Undo :func:ep_dispatch and restore the complete MoE head axis.

gather_sequence

gather_sequence(
    tensor: Tensor,
    split_sizes: list[int],
    group: Magi2ParallelGroup | None = None,
) -> Tensor

Gather uneven contiguous sequence shards in rank order.

get_magi2_ep_group

get_magi2_ep_group() -> Magi2ParallelGroup

Return the process group used for MAGI's MoE-head parallelism.

TP is the explicit MoE-head axis when it is larger than one. Otherwise the released SP-only layout overlaps head parallelism with Ulysses. Both groups are initialized and owned by vLLM-Omni; MAGI creates no ad-hoc process groups.

get_magi2_replica_group

get_magi2_replica_group(
    data_parallel_size: int,
) -> Magi2ParallelGroup

Return the complete TP x SP group for one data-parallel replica.

The pipeline uses this group for conditioning broadcasts and output-rank ownership. With DP=1 the diffusion world is exactly one TP x SP replica. With DP>1, MAGI currently requires TP=1, so the existing SP group is the complete per-replica group. CFG x DP is deliberately rejected by _validate_native_topology; if that combination is enabled in the future, this helper must select a group fixed to one DP replica and include its CFG ranks rather than returning the SP group alone.

get_magi2_tp_group

get_magi2_tp_group() -> Magi2ParallelGroup

Return vLLM's tensor-parallel group without an SP fallback.

get_magi2_ulysses_group

get_magi2_ulysses_group() -> Magi2ParallelGroup

Return vLLM-Omni's Ulysses subgroup, or a rank-local fallback.

MAGI-2 does not support Ring or AllGather-KV attention. The caller is expected to validate that ulysses_degree == sequence_parallel_size.

scatter_heads_gather_seqlen

scatter_heads_gather_seqlen(
    tensors: Iterable[Tensor],
    split_sizes: list[int],
    group: Magi2ParallelGroup | None = None,
) -> list[Tensor]

Inverse batched Ulysses exchange for Q/K/V.

Each input is [S_rank, world*H_i, D] and each output is [sum(S_r), H_i, D]. Fusing Q/K/V into one all-to-all preserves the reference communication ordering and avoids three independent collectives.

scatter_seqlen_gather_heads

scatter_seqlen_gather_heads(
    tensor: Tensor,
    split_sizes: list[int],
    group: Magi2ParallelGroup | None = None,
) -> Tensor

Ulysses [sum(S_r), H, D] -> [S_rank, world*H, D] exchange.

shard_sequence

shard_sequence(
    tensor: Tensor,
    split_sizes: list[int] | None = None,
    group: Magi2ParallelGroup | None = None,
) -> tuple[Tensor, list[int]]

Return this Ulysses rank's contiguous sequence slice.

Inputs are replicated at the pipeline boundary, matching the diffusion runner contract, so no collective is needed for dispatch.