[Spec] Support logprobs with DFlash (#33459)
This commit is contained in:
@@ -2442,7 +2442,11 @@ class Scheduler(
|
||||
self._maybe_namespace_elastic_radix_cache(req)
|
||||
|
||||
if self.spec_algorithm.is_dflash_family():
|
||||
error_msg = validate_dflash_request(req, self.enable_overlap)
|
||||
error_msg = (
|
||||
"DSpark speculative decoding does not support return_logprob yet."
|
||||
if self.spec_algorithm.is_dspark() and req.return_logprob
|
||||
else validate_dflash_request(req, self.enable_overlap)
|
||||
)
|
||||
if error_msg is not None:
|
||||
req.set_finish_with_abort(error_msg)
|
||||
self.init_req_max_new_tokens(req)
|
||||
|
||||
@@ -841,9 +841,6 @@ def build_dflash_verify_target_probs(
|
||||
|
||||
|
||||
def validate_dflash_request(req: Req, enable_overlap: bool) -> Optional[str]:
|
||||
if req.return_logprob:
|
||||
return "DFLASH speculative decoding does not support return_logprob yet."
|
||||
|
||||
if enable_overlap and req.return_hidden_states:
|
||||
return "DFLASH speculative decoding does not support return_hidden_states yet."
|
||||
|
||||
|
||||
@@ -17,6 +17,7 @@ from sglang.srt.configs.hybrid_arch import mambaish_config
|
||||
from sglang.srt.distributed import get_tp_group
|
||||
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.logprob_processor import compute_spec_v2_logprobs
|
||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||
from sglang.srt.managers.scheduler import GenerationBatchResult
|
||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||
@@ -1383,10 +1384,6 @@ class DFlashWorkerV2(BaseSpecWorker):
|
||||
on_publish=None,
|
||||
grammar_barrier=None,
|
||||
) -> GenerationBatchResult:
|
||||
if getattr(batch, "return_logprob", False):
|
||||
raise ValueError(
|
||||
"DFLASH speculative decoding does not support return_logprob yet."
|
||||
)
|
||||
self._validate_phase1_sampling_support(batch)
|
||||
|
||||
if batch.forward_mode.is_extend() or batch.is_extend_in_batch:
|
||||
@@ -1835,6 +1832,18 @@ class DFlashWorkerV2(BaseSpecWorker):
|
||||
# accept_len; recompute it from the forced commit_lens.
|
||||
new_seq_lens = None
|
||||
|
||||
if batch.return_logprob:
|
||||
output_indices = torch.arange(
|
||||
bs * block_size, dtype=torch.int64, device=device
|
||||
).view(bs, block_size)
|
||||
compute_spec_v2_logprobs(
|
||||
batch,
|
||||
logits_output,
|
||||
out_tokens.reshape(-1),
|
||||
output_indices,
|
||||
block_size - 1,
|
||||
)
|
||||
|
||||
if self._need_mamba_verify_commit:
|
||||
assert seq_lens_pre_verify is not None
|
||||
self._update_target_mamba_state_after_verify(
|
||||
|
||||
Reference in New Issue
Block a user