flashinfer.comm.decode_cp_a2a_lse_reduce

flashinfer.comm.decode_cp_a2a_lse_reduce(partial_o: Tensor, partial_lse: Tensor, workspace: Tensor, cp_rank: int, cp_size: int, lse_mode: Literal['base2', 'basee'] = 'base2') → Tensor

Fuse an NCCL LSA DCP A2A exchange with the LSE-weighted reduce.

Send, receive synchronization, and reduction execute in one cooperative CUDA kernel.

This is a collective operation. Every rank must invoke it in the same order and the same number of times, including CUDA graph replays. All ranks must belong to one NCCL LSA/NVLink domain.

Parameters:
  • partial_o (torch.Tensor) – [batch, heads, cp_size, head_dim] CUDA tensor (fp16 or bf16), or more generally [..., cp_size, head_dim]. partial_o[..., peer, :] is the slice destined for that CP rank.

  • partial_lse (torch.Tensor) – [batch, heads, cp_size] CUDA float32 tensor, or more generally [..., cp_size]. Its leading dimensions must match partial_o. heads is simply the number of heads present in the input; it may be local or total because this operation does not shard the head axis.

  • workspace (torch.Tensor) – Rendezvoused NCCL symmetric-memory tensor from decode_cp_a2a_lse_reduce_create_workspace(). Reuse it only from one ordered CUDA stream.

  • cp_rank (int) – This rank’s index in the CP group.

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

  • lse_mode (Literal["base2", "basee"]) – Logarithm base used by partial_lse. "base2" is the default produced by FlashInfer MLA; "basee" selects natural-log LSE.

Returns:

[batch, heads, head_dim], or more generally [..., head_dim], in the same dtype as partial_o.

Return type:

torch.Tensor