Skip to content

vllm.v1.attention.ops.pcp

_gather_prefill_cache_inputs(tensors, slot_mapping, num_decode_tokens, shard_decode_requests=False)

Gather PCP cache inputs while preserving replicated KV-cache state.

PCP-only execution shards decode requests across ranks. In that mode each rank must gather the other owners' decode KV as well as partitioned prefill KV so every PCP rank retains a complete cache replica. DCP execution keeps decode requests replicated and uses the legacy prefill-only gather.

Source code in vllm/v1/attention/ops/pcp.py
def _gather_prefill_cache_inputs(
    tensors: tuple[torch.Tensor, ...],
    slot_mapping: torch.Tensor,
    num_decode_tokens: int,
    shard_decode_requests: bool = False,
) -> tuple[tuple[torch.Tensor, ...], torch.Tensor]:
    """Gather PCP cache inputs while preserving replicated KV-cache state.

    PCP-only execution shards decode requests across ranks. In that mode each
    rank must gather the other owners' decode KV as well as partitioned prefill
    KV so every PCP rank retains a complete cache replica. DCP execution keeps
    decode requests replicated and uses the legacy prefill-only gather.
    """
    local_num_tokens = tensors[0].shape[0]
    assert all(tensor.shape[0] == local_num_tokens for tensor in tensors)
    assert 0 <= num_decode_tokens <= local_num_tokens

    # Replicated draft decodes use unexpanded slot mappings, even with DP padding.
    if slot_mapping.shape[0] <= local_num_tokens:
        return (
            tuple(tensor[:num_decode_tokens] for tensor in tensors),
            slot_mapping[:num_decode_tokens],
        )

    pcp_group = get_pcp_group()
    pcp_size = pcp_group.world_size
    gathered_slot_mapping = slot_mapping[: pcp_size * local_num_tokens]
    if shard_decode_requests:
        gathered_inputs = tuple(
            pcp_group.all_gather(tensor.contiguous(), dim=0) for tensor in tensors
        )
        return gathered_inputs, gathered_slot_mapping

    if num_decode_tokens == local_num_tokens:
        return tensors, slot_mapping[:num_decode_tokens]

    gathered_prefills = tuple(
        pcp_group.all_gather(tensor[num_decode_tokens:].contiguous(), dim=0)
        for tensor in tensors
    )
    if num_decode_tokens == 0:
        return gathered_prefills, gathered_slot_mapping

    cache_inputs = tuple(
        torch.cat((tensor[:num_decode_tokens], gathered_prefill), dim=0)
        for tensor, gathered_prefill in zip(tensors, gathered_prefills)
    )
    rank_slot_mappings = gathered_slot_mapping.view(pcp_size, local_num_tokens)
    cache_slot_mapping = torch.cat(
        (
            rank_slot_mappings[0, :num_decode_tokens],
            rank_slot_mappings[:, num_decode_tokens:].flatten(),
        )
    )
    return cache_inputs, cache_slot_mapping