flashinfer.comm.decode_cp_a2a_lse_reduce_create_workspace

flashinfer.comm.decode_cp_a2a_lse_reduce_create_workspace(max_tokens: int, local_heads: int, cp_size: int, head_dim: int, dtype: dtype, group: Any) → Tensor

Create and rendezvous the fused op’s NCCL symmetric workspace.

All ranks in group must call this function collectively. The group must fit in one NCCL load/store-accessible (LSA) NVLink domain; multi-node groups spanning LSA domains are not supported. Allocate one workspace per group and reuse it for every invocation and CUDA graph replay. A workspace may only be used from one ordered CUDA stream; allocate a workspace per concurrent stream.

Parameters:
  • max_tokens (int) – Upper bound on the token/batch dimension of later calls.

  • local_heads (int) – Number of heads this rank keeps after the reduce-scatter.

  • cp_size (int) – Context-parallel group size.

  • head_dim (int) – Elements per head.

  • dtype (torch.dtype) – Storage type of partial_o / output (fp16 or bf16).

  • group (torch.distributed.ProcessGroup or str) – Process group (or group name) for the CP team.

Returns:

A rendezvoused NCCL symmetric-memory uint8 tensor. Allocate it once before CUDA graph capture and reuse it.

Return type:

torch.Tensor