Skip to content

vllm_omni.diffusion.distributed.a2a_permute

Fused permute-free Ulysses all-to-all over NCCL symmetric memory.

JIT-compiles the CUDA kernel from pytorch/pytorch#178230 (all_to_all_permute) and exposes it as two functional custom ops:

ulysses_qkv_fwd(x, group_name, world_size)  # (B, S/p, H, D) -> (B, S, H/p, D)
ulysses_o_rev(y, group_name, world_size)     # (B, S, H/p, D) -> (B, S/p, H, D)

These replace the synchronous all_to_all_4D (permute + NCCL all_to_all_single) used by Ulysses SP. All symmetric-memory bookkeeping (a shape-keyed, rendezvoused buffer cache + the copy-in) lives inside the ops, so from torch.compile's point of view each op is an opaque input -> output function: no graph break, no visible buffer mutation, no Python control flow in the traced graph. A register_fake provides output metadata so Dynamo keeps the op in-graph.

logger module-attribute

logger = init_logger(__name__)

clear_a2a_permute_workspaces

clear_a2a_permute_workspaces() -> None

Release cached symmetric-memory workspaces during worker shutdown.

ensure_a2a_permute_available

ensure_a2a_permute_available() -> None

Build the extension during worker/model initialization, not a request.

ulysses_o_rev

ulysses_o_rev(
    y: Tensor, group_name: str, world_size: int
) -> Tensor

ulysses_qkv_fwd

ulysses_qkv_fwd(
    x: Tensor, group_name: str, world_size: int
) -> Tensor