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
groupmust 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
uint8tensor. Allocate it once before CUDA graph capture and reuse it.- Return type:
torch.Tensor