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 matchpartial_o.headsis 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 aspartial_o.- Return type:
torch.Tensor