From 36853b8ffc78e03501f9a1d9b9f0b5f04cb8797b Mon Sep 17 00:00:00 2001 From: Jason Mancuso <7891333+jvmncs@users.noreply.github.com> Date: Wed, 5 Aug 2026 15:37:12 -0400 Subject: [PATCH] [Spec] Support logprobs with DFlash (#33459) --- python/sglang/srt/managers/scheduler.py | 6 +++++- python/sglang/srt/speculative/dflash_utils.py | 3 --- .../sglang/srt/speculative/dflash_worker_v2.py | 17 +++++++++++++---- test/registered/spec/dflash/test_dflash.py | 7 ++----- 4 files changed, 20 insertions(+), 13 deletions(-) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 5480ddacd..be5b8a219 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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) diff --git a/python/sglang/srt/speculative/dflash_utils.py b/python/sglang/srt/speculative/dflash_utils.py index d7de9bf49..284ae5661 100644 --- a/python/sglang/srt/speculative/dflash_utils.py +++ b/python/sglang/srt/speculative/dflash_utils.py @@ -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." diff --git a/python/sglang/srt/speculative/dflash_worker_v2.py b/python/sglang/srt/speculative/dflash_worker_v2.py index 64ee3ea60..23fccf90d 100644 --- a/python/sglang/srt/speculative/dflash_worker_v2.py +++ b/python/sglang/srt/speculative/dflash_worker_v2.py @@ -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( diff --git a/test/registered/spec/dflash/test_dflash.py b/test/registered/spec/dflash/test_dflash.py index 44da2500f..fbc4077e8 100644 --- a/test/registered/spec/dflash/test_dflash.py +++ b/test/registered/spec/dflash/test_dflash.py @@ -12,7 +12,7 @@ from sglang.test.kits.radix_cache_server_kit import ( gen_radix_tree, run_radix_attention_test, ) -from sglang.test.kits.spec_server_kits import SpecGrammarKit +from sglang.test.kits.spec_server_kits import SpecGrammarKit, SpecLogprobKit from sglang.test.test_utils import ( DEFAULT_DRAFT_MODEL_DFLASH, DEFAULT_TARGET_MODEL_DFLASH, @@ -32,6 +32,7 @@ class TestDFlashServerBase( GSM8KMixin, JSONConstrainedMixin, SpecGrammarKit, + SpecLogprobKit, ): max_running_requests = 64 attention_backend = "triton" if is_hip() else "flashinfer" @@ -132,10 +133,6 @@ class TestDFlashServerBase( self.assertEqual(outputs[0], outputs[1]) assert self.process.poll() is None - @unittest.skip("DFLASH rejects return_logprob at admission") - def test_grammar_logprob_count_matches_completion_tokens(self): - pass - class TestDFlashServerPage256(TestDFlashServerBase): page_size = 256