class PCPManager:
"""MRV2 PC batch manager.
The model runner keeps the global scheduled batch. This manager rewrites only
the per-step InputBatch into rank-local DualChunkSwap rows and keeps the
global-batch view private to restore to the global batch shape before
sampling/postprocess.
"""
def __init__(
self,
pcp_world_size: int,
pcp_rank: int,
device: torch.device,
shard_decode_requests: bool,
max_num_reqs: int | None = None,
max_num_tokens: int | None = None,
block_tables: BlockTables | None = None,
dcp_world_size: int = 1,
dcp_rank: int = 0,
cp_interleave: int = 1,
) -> None:
self.pcp_world_size = pcp_world_size
self.pcp_rank = pcp_rank
self.device = device
self.dcp_world_size = dcp_world_size
self.dcp_rank = dcp_rank
self.cp_interleave = cp_interleave
self.shard_decode_requests = shard_decode_requests
self._global_batch: InputBatch | None = None
self._local_batch: InputBatch | None = None
self._local_gather_idx: torch.Tensor | None = None
self.draft_prefill_batch: InputBatch | None = None
self._block_tables = block_tables
self._hidden_restore_idx: torch.Tensor | None = None
self._padded_gather_idx: torch.Tensor | None = None
self._gathered_kv_write_mask: torch.Tensor | None = None
self._pad_slot_id = torch.tensor(PAD_SLOT_ID, dtype=torch.int64, device=device)
max_num_local_reqs = 2 * max_num_reqs if max_num_reqs is not None else None
self._input_buffers = (
InputBuffers(max_num_local_reqs, max_num_tokens, device)
if max_num_local_reqs is not None and max_num_tokens is not None
else None
)
self._local_block_tables: tuple[torch.Tensor, ...] | None
self._local_block_table_ptrs: torch.Tensor | None
if block_tables is not None and max_num_local_reqs is not None:
self._local_block_tables = tuple(
table.new_zeros((max_num_local_reqs, table.shape[1]))
for table in block_tables.input_block_tables
)
self._local_block_table_ptrs = torch.tensor(
[table.data_ptr() for table in self._local_block_tables],
dtype=torch.uint64,
device=device,
)
else:
self._local_block_tables = None
self._local_block_table_ptrs = None
num_kv_cache_groups = (
block_tables.num_kv_cache_groups if block_tables is not None else 0
)
self._global_batch_slot_mappings = (
torch.empty(
num_kv_cache_groups,
max_num_tokens,
dtype=torch.int64,
device=device,
)
if max_num_tokens is not None and num_kv_cache_groups > 0
else None
)
self._gathered_kv_slot_mappings = (
torch.empty(
num_kv_cache_groups,
max_num_tokens * pcp_world_size,
dtype=torch.int64,
device=device,
)
if max_num_tokens is not None and num_kv_cache_groups > 0
else None
)
@staticmethod
def validate_config(
vllm_config: VllmConfig,
supports_mm_inputs: bool,
) -> None:
parallel_config = vllm_config.parallel_config
model_config = vllm_config.model_config
pcp_size = parallel_config.prefill_context_parallel_size
if pcp_size <= 1:
return
if not model_config.use_mla:
raise NotImplementedError("MRV2 PCP currently supports MLA models only.")
if model_config.is_encoder_decoder:
raise NotImplementedError(
"MRV2 PCP does not support encoder-decoder models yet."
)
if supports_mm_inputs:
raise NotImplementedError("MRV2 PCP does not support MM inputs yet.")
if vllm_config.lora_config is not None:
raise NotImplementedError("MRV2 PCP does not support LoRA yet.")
speculative_config = vllm_config.speculative_config
if speculative_config is not None:
if speculative_config.use_dspark():
dcp_size = parallel_config.decode_context_parallel_size
if (
dcp_size not in (1, pcp_size)
and speculative_config.draft_model_config.use_mla
):
raise NotImplementedError(
"MRV2 PCP DSpark requires DCP=1 or DCP=PCP; got "
f"DCP={dcp_size}, PCP={pcp_size}."
)
elif (
speculative_config.method != "mtp"
or speculative_config.use_multi_module_mtp()
):
raise NotImplementedError(
"MRV2 PCP only supports DSpark or single-module MTP "
"speculative decoding."
)
cudagraph_mode = vllm_config.compilation_config.cudagraph_mode
is_sparse_mla = hasattr(model_config.hf_text_config, "index_topk")
if parallel_config.decode_context_parallel_size > 1 and not is_sparse_mla:
# Dense MLA prefill sizes its DCP KV gather from each rank's own
# chunk rows, so the ranks' collectives diverge (#53573).
raise NotImplementedError("MRV2 PCP + DCP supports sparse MLA models only.")
if (
is_sparse_mla
and parallel_config.decode_context_parallel_size == 1
and cudagraph_mode != CUDAGraphMode.NONE
):
raise NotImplementedError(
"MRV2 sparse MLA PCP does not support CUDA graphs yet. "
"Set -cc.cudagraph_mode=NONE."
)
if (
cudagraph_mode.has_full_cudagraphs()
and not cudagraph_mode.separate_routine()
):
raise NotImplementedError(
"MRV2 PCP supports full CUDA graphs for decode-only routines. "
"Use FULL_DECODE_ONLY, FULL_AND_PIECEWISE, PIECEWISE, or NONE."
)
if (
parallel_config.decode_context_parallel_size > 1
and parallel_config.dcp_comm_backend != "ag_rs"
):
raise NotImplementedError(
"MRV2 PCP + DCP requires dcp_comm_backend='ag_rs'; got "
f"'{parallel_config.dcp_comm_backend}'."
)
@staticmethod
def _reorder_segments(
segments: list[RankSegment], is_prefilling: np.ndarray
) -> list[RankSegment]:
"""Order this rank's rows decodes-first, then prefills, canonically."""
def sort_key(segment: RankSegment) -> tuple[bool, int, int]:
req_idx = segment.global_batch_req_idx
return (
bool(is_prefilling[req_idx]),
req_idx,
segment.global_batch_slice.start,
)
segments.sort(key=sort_key)
rank_offset = 0
for index, segment in enumerate(segments):
segments[index] = replace(
segment,
rank_local_batch_slice=slice(
rank_offset, rank_offset + segment.num_tokens
),
)
rank_offset += segment.num_tokens
return segments
def replicated_requests(
self, num_scheduled_tokens: np.ndarray, is_prefilling: np.ndarray
) -> np.ndarray:
"""Per global request, whether every PCP rank gets the whole query."""
num_chunks = 2 * self.pcp_world_size
query_lens = np.asarray(num_scheduled_tokens, dtype=np.int64)
replicated = ~np.asarray(is_prefilling, dtype=np.bool_)
if self.dcp_world_size > 1:
chunk_sizes = (query_lens + num_chunks - 1) // num_chunks
drops_a_chunk = (num_chunks - 1) * chunk_sizes >= query_lens
replicated |= drops_a_chunk
return replicated
def _iter_rank_chunks(
self,
rank: int,
num_scheduled_tokens: np.ndarray,
is_prefilling: np.ndarray,
) -> Iterator[tuple[int, int, int]]:
"""Yield ``(request index, query offset, length)`` for one PCP rank.
PCP=4 partitions each prefill into eight chunks:
full: | 0 | 1 | 2 | 3 | 4 | 5 | 6 | 7 |
rank 0: 0 7
rank 1: 1 6
rank 2: 2 5
rank 3: 3 4
Decodes, and prefills too short to fill all eight, are replicated instead.
"""
decode_ordinal = 0
num_chunks = 2 * self.pcp_world_size
replicated = self.replicated_requests(num_scheduled_tokens, is_prefilling)
for global_batch_req_idx, num_tokens in enumerate(num_scheduled_tokens):
query_len = int(num_tokens)
if query_len == 0:
continue
chunk_indices: tuple[int, ...]
if not replicated[global_batch_req_idx]:
chunk_size = (query_len + num_chunks - 1) // num_chunks
chunk_indices = (rank, num_chunks - 1 - rank)
elif self.shard_decode_requests:
chunk_size = query_len
# KV and hidden states are gathered back to every PCP rank, so
# decode ownership does not need to persist across steps. Use a
# compact decode-only ordinal to keep each step exactly
# balanced even when prefills and zero-token rows are present.
owner_rank = decode_ordinal % self.pcp_world_size
decode_ordinal += 1
chunk_indices = (0,) if rank == owner_rank else ()
else: # DCP requires decode queries on every participating rank.
chunk_size = query_len
chunk_indices = (0,)
for chunk_idx in chunk_indices:
chunk_offset = chunk_idx * chunk_size
chunk_len = min(chunk_size, query_len - chunk_offset)
if chunk_len <= 0:
continue
yield global_batch_req_idx, chunk_offset, chunk_len
def _get_rank_segments(
self,
rank: int,
num_scheduled_tokens: np.ndarray,
is_prefilling: np.ndarray,
query_start_loc_np: np.ndarray,
) -> list[RankSegment]:
rank_segments = []
rank_offset = 0
for global_batch_req_idx, chunk_offset, chunk_len in self._iter_rank_chunks(
rank, num_scheduled_tokens, is_prefilling
):
global_batch_start = int(query_start_loc_np[global_batch_req_idx])
chunk_start = global_batch_start + chunk_offset
rank_segments.append(
RankSegment(
global_batch_req_idx=global_batch_req_idx,
global_batch_slice=slice(chunk_start, chunk_start + chunk_len),
rank_local_batch_slice=slice(rank_offset, rank_offset + chunk_len),
)
)
rank_offset += chunk_len
return self._reorder_segments(rank_segments, is_prefilling)
def _build_batch_layout(
self,
num_scheduled_tokens: np.ndarray,
num_computed_tokens: np.ndarray,
is_prefilling: np.ndarray,
query_start_loc_np: np.ndarray,
padded_num_tokens: int | None = None,
) -> tuple[list[list[RankSegment]], list[int]]:
replicated = self.replicated_requests(num_scheduled_tokens, is_prefilling)
segments_by_rank = []
per_rank_num_tokens = []
for rank in range(self.pcp_world_size):
segments = self._get_rank_segments(
rank,
num_scheduled_tokens,
is_prefilling,
query_start_loc_np,
)
num_rank_tokens = sum(segment.num_tokens for segment in segments)
segments_by_rank.append(segments)
per_rank_num_tokens.append(num_rank_tokens)
# PCP=2 example:
# global batch: [A B C D E F G]
# rank 0 / rank 1: [A B G] / [C D E F]
# padded gathered: [A B G _ | C D E F]
# hidden_restore_idx: [0, 1, 4, 5, 6, 7, 2]
# padded_gather_idx: [0, 1, 6, 0, 2, 3, 4, 5]
# Therefore global = gathered[hidden_restore_idx] and
# padded_gathered = global[padded_gather_idx].
hidden_restore_idx = np.empty(int(query_start_loc_np[-1]), dtype=np.int64)
if padded_num_tokens is None:
padded_num_tokens = max(per_rank_num_tokens)
elif padded_num_tokens < max(per_rank_num_tokens):
raise ValueError(
"PCP padded token count is smaller than the largest rank-local "
f"batch: {padded_num_tokens} < {max(per_rank_num_tokens)}."
)
num_expanded_tokens = padded_num_tokens * self.pcp_world_size
padded_gather_idx = np.zeros(num_expanded_tokens, dtype=np.int64)
gathered_kv_write_mask = np.zeros(num_expanded_tokens, dtype=np.bool_)
for rank, segments in enumerate(segments_by_rank):
expanded_rank_offset = rank * padded_num_tokens
for segment in segments:
padded_gathered_slice = slice(
expanded_rank_offset + segment.rank_local_batch_slice.start,
expanded_rank_offset + segment.rank_local_batch_slice.stop,
)
padded_gather_idx[padded_gathered_slice] = np.arange(
segment.global_batch_slice.start,
segment.global_batch_slice.stop,
dtype=np.int64,
)
# Replicated rows contain identical KV, so rank 0 is the
# canonical writer. A sharded decode has exactly one owner and
# therefore needs no de-duplication here.
is_sharded_decode = (
not bool(is_prefilling[segment.global_batch_req_idx])
and self.shard_decode_requests
)
if (
replicated[segment.global_batch_req_idx]
and not is_sharded_decode
and rank != 0
):
continue
gathered_kv_write_mask[padded_gathered_slice] = True
hidden_restore_idx[segment.global_batch_slice] = np.arange(
padded_gathered_slice.start,
padded_gathered_slice.stop,
dtype=np.int64,
)
self._hidden_restore_idx = async_tensor_h2d(
hidden_restore_idx, device=self.device
)
self._padded_gather_idx = async_tensor_h2d(
padded_gather_idx, device=self.device
)
self._gathered_kv_write_mask = async_tensor_h2d(
gathered_kv_write_mask, device=self.device
)
return segments_by_rank, per_rank_num_tokens
def get_num_tokens_for_dispatch(
self,
num_scheduled_tokens: np.ndarray,
is_prefilling: np.ndarray,
) -> int:
"""Return the largest real rank-local batch before graph padding."""
return max(
sum(
chunk_len
for _, _, chunk_len in self._iter_rank_chunks(
rank, num_scheduled_tokens, is_prefilling
)
)
for rank in range(self.pcp_world_size)
)
@staticmethod
def _resolve_num_reqs_after_padding(
input_batch: InputBatch,
padded_num_reqs: int | None,
num_local_reqs: int,
) -> int:
if padded_num_reqs is None:
return num_local_reqs
if input_batch.has_prefill:
raise RuntimeError("PCP FULL graphs require a decode-only batch.")
assert padded_num_reqs >= num_local_reqs, (
"PCP graph request capacity must cover the rank-local batch: "
f"{padded_num_reqs} < {num_local_reqs}."
)
return padded_num_reqs
@property
def input_buffers(self) -> InputBuffers:
assert self._input_buffers is not None
return self._input_buffers
@property
def global_batch(self) -> InputBatch:
assert self._global_batch is not None
return self._global_batch
def partition_batch(
self, input_batch: InputBatch, batch_desc: "BatchExecutionDescriptor"
) -> InputBatch:
assert self._input_buffers is not None
input_buffers = self._input_buffers
global_batch = input_batch
self._global_batch = global_batch
num_scheduled_tokens = global_batch.num_scheduled_tokens
num_computed_tokens = global_batch.num_computed_tokens_np
is_prefilling = global_batch.is_prefilling_np
padded_num_tokens = None
padded_num_reqs = None
if batch_desc.cg_mode != CUDAGraphMode.NONE:
padded_num_tokens = batch_desc.num_tokens
if batch_desc.cg_mode == CUDAGraphMode.FULL:
padded_num_reqs = batch_desc.num_reqs
segments_by_rank, per_rank_num_tokens = self._build_batch_layout(
num_scheduled_tokens,
num_computed_tokens,
is_prefilling,
global_batch.query_start_loc_np,
padded_num_tokens=padded_num_tokens,
)
local_segments = segments_by_rank[self.pcp_rank]
if not local_segments:
local_segments = [
RankSegment(
global_batch_req_idx=0,
global_batch_slice=slice(0, 0),
rank_local_batch_slice=slice(0, 0),
)
]
num_local_reqs = len(local_segments)
num_reqs_after_padding = self._resolve_num_reqs_after_padding(
global_batch, padded_num_reqs, num_local_reqs
)
if num_reqs_after_padding > input_buffers.max_num_reqs:
raise RuntimeError(
"PCP padded local request count exceeds the MRV2 input buffer size: "
f"{num_reqs_after_padding} > {input_buffers.max_num_reqs}."
)
local_to_global_batch_req_idx_np = np.fromiter(
(segment.global_batch_req_idx for segment in local_segments),
dtype=np.intp,
count=num_local_reqs,
)
local_start_pos_np = np.fromiter(
(
num_computed_tokens[segment.global_batch_req_idx]
+ segment.global_batch_slice.start
- global_batch.query_start_loc_np[segment.global_batch_req_idx]
for segment in local_segments
),
dtype=np.int32,
count=num_local_reqs,
)
local_num_scheduled_tokens = np.fromiter(
(segment.num_tokens for segment in local_segments),
dtype=np.int32,
count=num_local_reqs,
)
local_to_global_req_idx_np = global_batch.idx_mapping_np[
local_to_global_batch_req_idx_np
]
local_req_ids = [
global_batch.req_ids[global_batch_req_idx]
for global_batch_req_idx in local_to_global_batch_req_idx_np
]
num_local_tokens = int(local_num_scheduled_tokens.sum())
num_local_tokens_padded = (
max(per_rank_num_tokens) if padded_num_tokens is None else padded_num_tokens
)
fresh_prefills = int(
np.count_nonzero(is_prefilling & (num_computed_tokens == 0))
)
continued_prefills = int(
np.count_nonzero(is_prefilling & (num_computed_tokens > 0))
)
logger.debug(
"PCP batch: rank=%d global_batch_reqs=%d fresh_prefills=%d "
"continued_prefills=%d decodes=%d local_reqs=%d "
"local_tokens=%d per_rank_tokens=%s",
self.pcp_rank,
global_batch.num_reqs,
fresh_prefills,
continued_prefills,
global_batch.num_reqs - fresh_prefills - continued_prefills,
num_local_reqs,
num_local_tokens,
per_rank_num_tokens,
)
if num_local_tokens_padded > input_buffers.max_num_tokens:
raise RuntimeError(
"PCP local token count exceeds the MRV2 input buffer size: "
f"{num_local_tokens_padded} > {input_buffers.max_num_tokens}."
)
rank_token_start = self.pcp_rank * num_local_tokens_padded
assert self._padded_gather_idx is not None
local_gather_idx = self._padded_gather_idx[
rank_token_start : rank_token_start + num_local_tokens_padded
]
self._local_gather_idx = local_gather_idx
torch.index_select(
global_batch.input_ids,
0,
local_gather_idx,
out=input_buffers.input_ids[:num_local_tokens_padded],
)
# Keep the GPU request-state cursor materialized by prepare_inputs().
# The CPU cursor can lag after speculative rejection.
torch.index_select(
global_batch.positions,
0,
local_gather_idx,
out=input_buffers.positions[:num_local_tokens_padded],
)
local_query_start_loc_np = np.empty(
input_buffers.max_num_reqs + 1, dtype=np.int32
)
local_query_start_loc_np[0] = 0
local_query_start_loc_out = local_query_start_loc_np[1 : num_local_reqs + 1]
np.cumsum(local_num_scheduled_tokens, out=local_query_start_loc_out)
local_query_start_loc_np[num_local_reqs + 1 :] = num_local_tokens
async_tensor_h2d(local_query_start_loc_np, out=input_buffers.query_start_loc)
local_query_start_loc = input_buffers.query_start_loc[
: num_reqs_after_padding + 1
]
local_to_global_req_idx = async_tensor_h2d(
local_to_global_req_idx_np, device=self.device
)
seq_lens = input_buffers.seq_lens[:num_reqs_after_padding]
real_seq_lens = seq_lens[:num_local_reqs]
if num_local_tokens > 0:
local_end_positions = torch.index_select(
input_buffers.positions,
0,
local_query_start_loc[1 : num_local_reqs + 1] - 1,
)
real_seq_lens.copy_(local_end_positions + 1)
else:
real_seq_lens.zero_()
seq_lens[num_local_reqs:].zero_()
is_padding = input_buffers.is_padding[:num_local_tokens_padded]
is_padding[:num_local_tokens].fill_(False)
is_padding[num_local_tokens:].fill_(True)
if num_local_tokens_padded > num_local_tokens:
input_buffers.input_ids[:num_local_tokens_padded].masked_fill_(
is_padding, 0
)
input_buffers.positions[:num_local_tokens_padded].masked_fill_(
is_padding, 0
)
total_num_logits = num_local_reqs if num_local_tokens > 0 else 0
if total_num_logits > 0:
cu_num_logits_np = np.arange(num_local_reqs + 1, dtype=np.int32)
cu_num_logits = torch.arange(
num_local_reqs + 1, device=self.device, dtype=torch.int32
)
else:
cu_num_logits_np = np.zeros(num_local_reqs + 1, dtype=np.int32)
cu_num_logits = torch.zeros(
num_local_reqs + 1, device=self.device, dtype=torch.int32
)
# Local logits are never sampled. The complete hidden-state tensor is
# restored first and sampled with the untouched global InputBatch.
logits_indices = local_query_start_loc[1:] - 1
local_prefill_len_np = global_batch.prefill_len_np[
local_to_global_batch_req_idx_np
]
local_num_computed_prefill_tokens_np = np.minimum(
local_start_pos_np, local_prefill_len_np
)
real_local_is_prefilling_np = (
local_num_computed_prefill_tokens_np < local_prefill_len_np
)
local_is_prefilling_np = np.zeros(num_reqs_after_padding, dtype=np.bool_)
local_is_prefilling_np[:num_local_reqs] = real_local_is_prefilling_np
local_has_prefill = bool(local_is_prefilling_np.any())
seq_lens_cpu_upper_bound_np = np.zeros(num_reqs_after_padding, dtype=np.int32)
seq_lens_cpu_upper_bound_np[:num_local_reqs] = (
local_start_pos_np + local_num_scheduled_tokens
)
dcp_local_seq_lens_cpu_upper_bound = None
if self.dcp_world_size > 1:
# The largest DCP shard of each row's whole request, identical on
# every PCP rank: the sparse backends pad their KV gather to it.
request_seq_lens = (num_computed_tokens + num_scheduled_tokens)[
local_to_global_batch_req_idx_np
]
real_dcp_local_seq_lens_cpu_upper_bound = get_dcp_local_seq_lens(
torch.from_numpy(request_seq_lens.astype(np.int32)),
self.dcp_world_size,
0,
self.cp_interleave,
)
dcp_local_seq_lens_cpu_upper_bound = torch.zeros(
num_reqs_after_padding, dtype=torch.int32
)
dcp_local_seq_lens_cpu_upper_bound[:num_local_reqs].copy_(
real_dcp_local_seq_lens_cpu_upper_bound
)
self._local_batch = replace(
input_batch,
req_ids=local_req_ids,
num_reqs=num_local_reqs,
num_reqs_after_padding=num_reqs_after_padding,
idx_mapping=local_to_global_req_idx,
idx_mapping_np=local_to_global_req_idx_np,
expanded_idx_mapping=local_to_global_req_idx,
expanded_local_pos=torch.zeros(
num_local_reqs, dtype=torch.int32, device=self.device
),
num_scheduled_tokens=local_num_scheduled_tokens,
num_tokens=num_local_tokens,
num_tokens_after_padding=num_local_tokens_padded,
num_draft_tokens=0,
num_draft_tokens_per_req=None,
query_start_loc=local_query_start_loc,
query_start_loc_np=local_query_start_loc_np[: num_reqs_after_padding + 1],
seq_lens=seq_lens,
seq_lens_cpu_upper_bound=torch.from_numpy(seq_lens_cpu_upper_bound_np),
dcp_local_seq_lens=None,
dcp_local_seq_lens_cpu_upper_bound=dcp_local_seq_lens_cpu_upper_bound,
num_computed_tokens_np=local_start_pos_np,
prefill_len_np=local_prefill_len_np,
num_computed_prefill_tokens_np=local_num_computed_prefill_tokens_np,
is_prefilling_np=local_is_prefilling_np,
max_seq_len_np=None,
has_prefill=local_has_prefill,
decode_graph_eligible=not local_has_prefill,
prefill_runs_as_decode_np=None,
input_ids=input_buffers.input_ids[:num_local_tokens_padded],
positions=input_buffers.positions[:num_local_tokens_padded],
is_padding=is_padding,
logits_indices=logits_indices,
cu_num_logits=cu_num_logits,
cu_num_logits_np=cu_num_logits_np,
prompt_lens=None,
)
return self._local_batch
def prepare_inputs_to_capture(self, input_batch: InputBatch) -> InputBatch:
"""Stage a capture or dummy batch in persistent PCP input buffers."""
input_buffers = self.input_buffers
num_reqs = input_batch.num_reqs_after_padding
num_tokens = input_batch.num_tokens_after_padding
input_batch = replace(
input_batch,
input_ids=input_buffers.input_ids[:num_tokens].copy_(input_batch.input_ids),
positions=input_buffers.positions[:num_tokens].copy_(input_batch.positions),
is_padding=input_buffers.is_padding[:num_tokens].copy_(
input_batch.is_padding
),
query_start_loc=input_buffers.query_start_loc[: num_reqs + 1].copy_(
input_batch.query_start_loc
),
seq_lens=input_buffers.seq_lens[:num_reqs].copy_(input_batch.seq_lens),
)
return input_batch
def get_dummy_block_tables(self, num_reqs: int) -> tuple[torch.Tensor, ...]:
assert self._local_block_tables is not None
return tuple(
block_table[:num_reqs].zero_() for block_table in self._local_block_tables
)
def prepare_attn(
self, input_batch: InputBatch
) -> tuple[tuple[torch.Tensor, ...], torch.Tensor]:
assert self._block_tables is not None
assert self._local_block_tables is not None
assert self._local_block_table_ptrs is not None
block_tables = self._block_tables.gather_block_tables(
input_batch.idx_mapping,
input_batch.num_reqs_after_padding,
out=self._local_block_tables,
out_ptrs=self._local_block_table_ptrs,
)
slot_mappings = self.prepare_slot_mappings()
return block_tables, slot_mappings
def prepare_slot_mappings(self) -> torch.Tensor:
assert self._block_tables is not None
assert self._global_batch_slot_mappings is not None
assert self._global_batch is not None
global_batch = self._global_batch
global_batch_slot_mappings = self._block_tables.compute_slot_mappings(
global_batch.idx_mapping,
global_batch.query_start_loc,
global_batch.positions,
global_batch.num_tokens,
out=self._global_batch_slot_mappings,
)
return self._convert_to_gathered_slot_mappings(global_batch_slot_mappings)
def get_dummy_slot_mappings(self, num_tokens: int) -> torch.Tensor:
assert self._gathered_kv_slot_mappings is not None
self._gathered_kv_slot_mappings.fill_(PAD_SLOT_ID)
return self._gathered_kv_slot_mappings[:, : num_tokens * self.pcp_world_size]
def _convert_to_gathered_slot_mappings(
self, global_batch_slot_mappings: torch.Tensor
) -> torch.Tensor:
assert self._padded_gather_idx is not None
assert self._gathered_kv_write_mask is not None
padded_gather_idx = self._padded_gather_idx
num_expanded_tokens = padded_gather_idx.shape[0]
if self._gathered_kv_slot_mappings is None:
self._gathered_kv_slot_mappings = global_batch_slot_mappings.new_empty(
global_batch_slot_mappings.shape[0], num_expanded_tokens
)
gathered_kv_slot_mappings = self._gathered_kv_slot_mappings[
:, :num_expanded_tokens
]
torch.index_select(
global_batch_slot_mappings,
1,
padded_gather_idx,
out=gathered_kv_slot_mappings,
)
torch.where(
self._gathered_kv_write_mask.unsqueeze(0),
gathered_kv_slot_mappings,
self._pad_slot_id,
out=gathered_kv_slot_mappings,
)
return gathered_kv_slot_mappings
def restore_hidden_states(self, hidden_states: torch.Tensor) -> torch.Tensor:
if self._hidden_restore_idx is None:
return hidden_states
gathered = get_pcp_group().all_gather(hidden_states, dim=0)
return gathered[self._hidden_restore_idx]
def get_draft_input_buffers(
self, input_buffers: InputBuffers
) -> InputBatch | InputBuffers:
return self.draft_prefill_batch or input_buffers
def prepare_draft_prefill(
self, input_batch: InputBatch, input_ids: torch.Tensor
) -> None:
self.draft_prefill_batch = None
if input_batch is not self._global_batch or self._local_batch is None:
return
local_batch = self._local_batch
assert self._local_gather_idx is not None
num_local_tokens = self._local_gather_idx.shape[0]
torch.index_select(
input_ids,
0,
self._local_gather_idx,
out=local_batch.input_ids[:num_local_tokens],
)
self.draft_prefill_batch = local_batch
def restore_draft_prefill(
self,
last_hidden_states: torch.Tensor,
hidden_states: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
if self.draft_prefill_batch is None:
return last_hidden_states, hidden_states
local_last_hidden_states = last_hidden_states
last_hidden_states = self.restore_hidden_states(local_last_hidden_states)
hidden_states = (
last_hidden_states
if local_last_hidden_states is hidden_states
else self.restore_hidden_states(hidden_states)
)
self.draft_prefill_batch = None
return last_hidden_states, hidden_states
def restore_for_sampling(
self, hidden_states: torch.Tensor, aux_hidden_states: list[torch.Tensor] | None
) -> tuple[torch.Tensor, list[torch.Tensor] | None, InputBatch]:
assert self._global_batch is not None
hidden_states = self.restore_hidden_states(hidden_states)
if aux_hidden_states is not None:
aux_hidden_states = [
self.restore_hidden_states(states) for states in aux_hidden_states
]
return hidden_states, aux_hidden_states, self._global_batch