[PD] Support logprob & Add failure test (#6558)
This commit is contained in:
@@ -5,7 +5,7 @@ from typing import TYPE_CHECKING
|
||||
|
||||
import torch
|
||||
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode
|
||||
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode, ForwardMode
|
||||
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -76,6 +76,11 @@ class ScheduleBatchDisaggregationDecodeMixin:
|
||||
self.seq_lens = torch.tensor(seq_lens, dtype=torch.int64, device=self.device)
|
||||
self.out_cache_loc = out_cache_loc
|
||||
self.seq_lens_sum = sum(seq_lens)
|
||||
|
||||
if self.return_logprob:
|
||||
self.top_logprobs_nums = [r.top_logprobs_num for r in reqs]
|
||||
self.token_ids_logprobs = [r.token_ids_logprob for r in reqs]
|
||||
|
||||
self.extend_num_tokens = extend_num_tokens
|
||||
self.prefix_lens = [len(r.prefix_indices) for r in reqs]
|
||||
self.extend_lens = [r.extend_input_len for r in reqs]
|
||||
@@ -94,12 +99,41 @@ class ScheduleBatchDisaggregationDecodeMixin:
|
||||
"""Assign the buffered last input id to schedule batch"""
|
||||
self.output_ids = []
|
||||
for req in self.reqs:
|
||||
if req.output_ids and len(req.output_ids) > 0:
|
||||
# resumed retracted req
|
||||
self.output_ids.append(req.output_ids[-1])
|
||||
else:
|
||||
assert req.transferred_output_id is not None
|
||||
req.output_ids.append(req.transferred_output_id)
|
||||
self.output_ids.append(req.transferred_output_id)
|
||||
self.output_ids.append(req.output_ids[-1])
|
||||
self.tree_cache.cache_unfinished_req(req)
|
||||
self.output_ids = torch.tensor(self.output_ids, device=self.device)
|
||||
|
||||
# Simulate the eagle run. We add mock data to hidden states for the
|
||||
# ease of implementation now meaning the first token will have acc rate
|
||||
# of 0.
|
||||
if not self.spec_algorithm.is_none():
|
||||
|
||||
b = len(self.reqs)
|
||||
topk_p = torch.arange(
|
||||
b * server_args.speculative_eagle_topk,
|
||||
0,
|
||||
-1,
|
||||
device=self.device,
|
||||
dtype=torch.float32,
|
||||
)
|
||||
topk_p = topk_p.reshape(b, server_args.speculative_eagle_topk)
|
||||
topk_p /= b * server_args.speculative_eagle_topk
|
||||
topk_index = torch.arange(
|
||||
b * server_args.speculative_eagle_topk, device=self.device
|
||||
)
|
||||
topk_index = topk_index.reshape(b, server_args.speculative_eagle_topk)
|
||||
|
||||
# local import to avoid circular import
|
||||
from sglang.srt.speculative.eagle_utils import EagleDraftInput
|
||||
|
||||
spec_info = EagleDraftInput(
|
||||
topk_p=topk_p,
|
||||
topk_index=topk_index,
|
||||
hidden_states=torch.ones(
|
||||
(b, model_config.hidden_size), device=self.device
|
||||
),
|
||||
verified_id=self.output_ids,
|
||||
)
|
||||
spec_info.prepare_for_extend(self)
|
||||
spec_info.capture_hidden_mode = CaptureHiddenMode.LAST
|
||||
self.spec_info = spec_info
|
||||
|
||||
Reference in New Issue
Block a user