[PD] Support logprob & Add failure test (#6558)
This commit is contained in:
@@ -36,6 +36,7 @@ from sglang.srt.disaggregation.utils import (
|
||||
DisaggregationMode,
|
||||
FakeBootstrapHost,
|
||||
KVClassType,
|
||||
MetadataBuffers,
|
||||
ReqToMetadataIdxAllocator,
|
||||
TransferBackend,
|
||||
get_kv_class,
|
||||
@@ -78,8 +79,7 @@ class DecodePreallocQueue:
|
||||
token_to_kv_pool_allocator: TokenToKVPoolAllocator,
|
||||
draft_token_to_kv_pool: Optional[KVCache],
|
||||
req_to_metadata_buffer_idx_allocator: ReqToMetadataIdxAllocator,
|
||||
metadata_buffers: List[torch.Tensor],
|
||||
aux_dtype: torch.dtype,
|
||||
metadata_buffers: MetadataBuffers,
|
||||
scheduler: Scheduler,
|
||||
transfer_queue: DecodeTransferQueue,
|
||||
tree_cache: BasePrefixCache,
|
||||
@@ -94,7 +94,6 @@ class DecodePreallocQueue:
|
||||
self.token_to_kv_pool = token_to_kv_pool_allocator.get_kvcache()
|
||||
self.draft_token_to_kv_pool = draft_token_to_kv_pool
|
||||
self.is_mla_backend = is_mla_backend(self.token_to_kv_pool)
|
||||
self.aux_dtype = aux_dtype
|
||||
self.metadata_buffers = metadata_buffers
|
||||
self.req_to_metadata_buffer_idx_allocator = req_to_metadata_buffer_idx_allocator
|
||||
self.scheduler = scheduler
|
||||
@@ -133,15 +132,9 @@ class DecodePreallocQueue:
|
||||
kv_args.kv_data_lens = kv_data_lens
|
||||
kv_args.kv_item_lens = kv_item_lens
|
||||
|
||||
kv_args.aux_data_ptrs = [
|
||||
output_id_tensor.data_ptr() for output_id_tensor in self.metadata_buffers
|
||||
]
|
||||
kv_args.aux_data_lens = [
|
||||
metadata_buffer.nbytes for metadata_buffer in self.metadata_buffers
|
||||
]
|
||||
kv_args.aux_item_lens = [
|
||||
metadata_buffer[0].nbytes for metadata_buffer in self.metadata_buffers
|
||||
]
|
||||
kv_args.aux_data_ptrs, kv_args.aux_data_lens, kv_args.aux_item_lens = (
|
||||
self.metadata_buffers.get_buf_infos()
|
||||
)
|
||||
kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device
|
||||
kv_args.gpu_id = self.scheduler.gpu_id
|
||||
kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
|
||||
@@ -211,7 +204,18 @@ class DecodePreallocQueue:
|
||||
indices_to_remove = set()
|
||||
allocatable_tokens = self._allocatable_tokens()
|
||||
|
||||
# First, remove all failed requests from the queue
|
||||
for i, decode_req in enumerate(self.queue):
|
||||
if isinstance(decode_req.req.finished_reason, FINISH_ABORT):
|
||||
self.scheduler.stream_output(
|
||||
[decode_req.req], decode_req.req.return_logprob
|
||||
)
|
||||
indices_to_remove.add(i)
|
||||
|
||||
for i, decode_req in enumerate(self.queue):
|
||||
if i in indices_to_remove:
|
||||
continue
|
||||
|
||||
if not decode_req.waiting_for_input:
|
||||
continue
|
||||
|
||||
@@ -331,7 +335,7 @@ class DecodeTransferQueue:
|
||||
self,
|
||||
gloo_group: ProcessGroup,
|
||||
req_to_metadata_buffer_idx_allocator: ReqToMetadataIdxAllocator,
|
||||
metadata_buffers: torch.Tensor,
|
||||
metadata_buffers: MetadataBuffers,
|
||||
scheduler: Scheduler,
|
||||
tree_cache: BasePrefixCache,
|
||||
):
|
||||
@@ -342,11 +346,11 @@ class DecodeTransferQueue:
|
||||
self.scheduler = scheduler
|
||||
self.tree_cache = tree_cache
|
||||
|
||||
def add(self, req_conn: DecodeRequest) -> None:
|
||||
self.queue.append(req_conn)
|
||||
def add(self, decode_req: DecodeRequest) -> None:
|
||||
self.queue.append(decode_req)
|
||||
|
||||
def extend(self, req_conns) -> None:
|
||||
self.queue.extend(req_conns)
|
||||
def extend(self, decode_reqs: List[DecodeRequest]) -> None:
|
||||
self.queue.extend(decode_reqs)
|
||||
|
||||
def pop_transferred(self) -> List[DecodeRequest]:
|
||||
if not self.queue:
|
||||
@@ -356,14 +360,6 @@ class DecodeTransferQueue:
|
||||
[decode_req.kv_receiver for decode_req in self.queue], self.gloo_group
|
||||
)
|
||||
|
||||
# First, remove all failed requests from the queue
|
||||
for i, decode_req in enumerate(self.queue):
|
||||
if isinstance(decode_req.req.finished_reason, FINISH_ABORT):
|
||||
self.scheduler.stream_output(
|
||||
[decode_req.req], decode_req.req.return_logprob
|
||||
)
|
||||
indices_to_remove.add(i)
|
||||
|
||||
transferred_reqs = []
|
||||
indices_to_remove = set()
|
||||
for i, (decode_req, poll) in enumerate(zip(self.queue, polls)):
|
||||
@@ -387,16 +383,37 @@ class DecodeTransferQueue:
|
||||
indices_to_remove.add(i)
|
||||
continue
|
||||
elif poll == KVPoll.Success:
|
||||
# pop and push it to waiting queue
|
||||
|
||||
idx = decode_req.metadata_buffer_index
|
||||
assert len(decode_req.req.output_ids) == 0
|
||||
output_id_buffer = self.metadata_buffers[0]
|
||||
# the last dimension is padded by the same values.
|
||||
output_id = output_id_buffer[idx][0].item()
|
||||
assert len(decode_req.req.output_ids) == 0
|
||||
assert decode_req.req.transferred_output_id is None
|
||||
decode_req.req.transferred_output_id = output_id
|
||||
transferred_reqs.append(decode_req)
|
||||
(
|
||||
output_id,
|
||||
output_token_logprobs_val,
|
||||
output_token_logprobs_idx,
|
||||
output_top_logprobs_val,
|
||||
output_top_logprobs_idx,
|
||||
) = self.metadata_buffers.get_buf(idx)
|
||||
|
||||
decode_req.req.output_ids.append(output_id[0].item())
|
||||
|
||||
if decode_req.req.return_logprob:
|
||||
decode_req.req.output_token_logprobs_val.append(
|
||||
output_token_logprobs_val[0].item()
|
||||
)
|
||||
decode_req.req.output_token_logprobs_idx.append(
|
||||
output_token_logprobs_idx[0].item()
|
||||
)
|
||||
decode_req.req.output_top_logprobs_val.append(
|
||||
output_top_logprobs_val[
|
||||
: decode_req.req.top_logprobs_num
|
||||
].tolist()
|
||||
)
|
||||
decode_req.req.output_top_logprobs_idx.append(
|
||||
output_top_logprobs_idx[
|
||||
: decode_req.req.top_logprobs_num
|
||||
].tolist()
|
||||
)
|
||||
|
||||
transferred_reqs.append(decode_req.req)
|
||||
indices_to_remove.add(i)
|
||||
elif poll in [
|
||||
KVPoll.Bootstrapping,
|
||||
@@ -451,7 +468,9 @@ class SchedulerDisaggregationDecodeMixin:
|
||||
# Generate fake extend output.
|
||||
if batch.forward_mode.is_extend():
|
||||
# Note: Logprobs should be handled on the prefill engine.
|
||||
self.stream_output(batch.reqs, False)
|
||||
self.stream_output(
|
||||
batch.reqs, any(req.return_logprob for req in batch.reqs)
|
||||
)
|
||||
if prepare_dp_attn_flag:
|
||||
self._prepare_idle_batch_and_run(None)
|
||||
else:
|
||||
@@ -497,7 +516,9 @@ class SchedulerDisaggregationDecodeMixin:
|
||||
# Generate fake extend output.
|
||||
if batch.forward_mode.is_extend():
|
||||
# Note: Logprobs should be handled on the prefill engine.
|
||||
self.stream_output(batch.reqs, False)
|
||||
self.stream_output(
|
||||
batch.reqs, any(req.return_logprob for req in batch.reqs)
|
||||
)
|
||||
if prepare_dp_attn_flag:
|
||||
batch_, result = self._prepare_idle_batch_and_run(
|
||||
None, delay_process=True
|
||||
@@ -618,15 +639,8 @@ class SchedulerDisaggregationDecodeMixin:
|
||||
|
||||
def process_decode_queue(self: Scheduler):
|
||||
req_conns = self.disagg_decode_prealloc_queue.pop_preallocated()
|
||||
|
||||
def _num_pre_alloc(req):
|
||||
return len(req.req.origin_input_ids) + max(len(req.req.output_ids) - 1, 0)
|
||||
|
||||
self.num_tokens_pre_allocated += sum(_num_pre_alloc(req) for req in req_conns)
|
||||
self.disagg_decode_transfer_queue.extend(req_conns)
|
||||
alloc_reqs = (
|
||||
self.disagg_decode_transfer_queue.pop_transferred()
|
||||
) # the requests which kv has arrived
|
||||
self.num_tokens_pre_allocated -= sum(_num_pre_alloc(req) for req in alloc_reqs)
|
||||
|
||||
self.waiting_queue.extend([req.req for req in alloc_reqs])
|
||||
self.waiting_queue.extend(alloc_reqs)
|
||||
|
||||
Reference in New Issue
Block a user