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.