flashinfer.comm.ulysses_a2a

flashinfer.comm.ulysses_a2a(fa: int, inp: Tensor, out: Tensor, B: int, S_local: int, H: int, D: int, mode: int) None

Fused-transpose Ulysses all-to-all.

Note

Advanced / internal API. Prefer UlyssesCommunicator.scatter_heads() (mode == 0) and UlyssesCommunicator.gather_heads() (mode == 1), which derive the geometry from the tensor shapes and validate operands.

fa is the opaque backend handle returned by init_ulysses_a2a() (a C++ UlyssesA2A* reinterpreted as an integer fptr_t); it selects the all-to-all context to run on and is not a device pointer.

The result for this rank is written into out (bit-identical to the equivalent NCCL all-to-all followed by the layout permutation).

mode == 0 (input a2a): inp [B, S_local, H, D] -> out [B, S_global, H_local, D]

mode == 1 (output a2a): inp [B, S_global, H_local, D] -> out [B, S_local, H, D]

where H is the global head count, H_local = H // world_size and S_global = S_local * world_size. Both tensors must be contiguous CUDA tensors of the same dtype (float32/float16/bfloat16). All ranks must call with consistent geometry in the same order; a mismatch is a collective failure (hang or corruption), as with any collective.