From 4e6ce8e4ab69a2ff0edd7a636a4cf3efb811a99d Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Sat, 27 Jun 2026 19:34:26 -0700 Subject: [PATCH] [Spec] Capture DFLASH draft greedy sampling inside the draft decode cuda graph (#29395) --- .../runner/decode_cuda_graph_runner.py | 18 ++- .../srt/speculative/dflash_worker_v2.py | 103 ++++++++++++++++-- 2 files changed, 109 insertions(+), 12 deletions(-) diff --git a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py index 19eded792..bfbe94774 100644 --- a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py @@ -802,12 +802,28 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): ): kwargs["input_embeds"] = self.buffers.input_embeds[:num_tokens] - return forward( + out = forward( forward_batch.input_ids, forward_batch.positions, forward_batch, **kwargs, ) + dflash_sampler = getattr( + self.model_runner, "dflash_draft_sampler", None + ) + if dflash_sampler is not None: + # Must be captured here, or replay leaves a stale output buffer + # the worker would read as valid tokens -- fail loudly instead. + if ( + not isinstance(out, LogitsProcessorOutput) + or out.hidden_states is None + ): + raise RuntimeError( + "DFLASH draft sampler set but the draft forward has no " + "hidden_states to capture into the graph." + ) + dflash_sampler(out.hidden_states) + return out self.deepep_adapter.capture(is_extend_in_batch=False) canary_ctx = ( diff --git a/python/sglang/srt/speculative/dflash_worker_v2.py b/python/sglang/srt/speculative/dflash_worker_v2.py index 132719818..6e1ec178c 100644 --- a/python/sglang/srt/speculative/dflash_worker_v2.py +++ b/python/sglang/srt/speculative/dflash_worker_v2.py @@ -59,6 +59,40 @@ def _get_fused_kv_materialize_helper(): return _FusedKVMaterializeHelper +class _DflashDraftSampler: + """Capture-safe greedy argmax over the target LM head, run inside the draft + cuda graph so the draft sampling is captured and counted in fwd_occupancy. + DFLASH's draft has no head of its own; it borrows the target `lm_head`. + tp=1 / no-added-vocab only; TP>1 stays eager in the worker. + """ + + def __init__(self, *, weight, block_size, num_org, org_vocab_start, max_bs): + self.weight = weight + self.block_size = int(block_size) + self.num_org = int(num_org) + self.org_vocab_start = int(org_vocab_start) + # Proposed draft tokens: written in-graph, read by the worker after replay. + self.out = torch.empty( + (int(max_bs) * (self.block_size - 1),), + dtype=torch.int64, + device=weight.device, + ) + + def __call__(self, hidden_states): + # draft tokens are block positions 1: (pos 0 is the seeded bonus token) + bs = hidden_states.shape[0] // self.block_size + hs = hidden_states.view(bs, self.block_size, -1)[:, 1:, :].reshape( + -1, hidden_states.shape[-1] + ) + if hs.dtype != self.weight.dtype: + hs = hs.to(self.weight.dtype) + logits = torch.matmul(hs, self.weight[: self.num_org].T) + tokens = torch.argmax(logits, dim=-1).to(torch.long) + if self.org_vocab_start: + tokens += self.org_vocab_start + self.out[: tokens.shape[0]].copy_(tokens) + + class DFlashWorkerV2(BaseSpecWorker): """DFLASH speculative decoding worker (spec-v2). @@ -158,6 +192,7 @@ class DFlashWorkerV2(BaseSpecWorker): ) set_global_server_args_for_scheduler(saved_server_args) self.draft_model_runner = self._draft_worker.model_runner + self._draft_sampler = None # Keep the same alias that other spec-v2 workers expose. self._draft_worker.draft_runner = self.draft_model_runner self.draft_model = self.draft_model_runner.model @@ -305,10 +340,50 @@ class DFlashWorkerV2(BaseSpecWorker): "memory is available after target backend initialization.", available_mem, ) + if capture_decode_cuda_graph: + # Must run before capture so the draft graph folds the head in. + self._draft_sampler = self._maybe_build_draft_sampler() + self.draft_model_runner.dflash_draft_sampler = self._draft_sampler self._draft_worker.init_cuda_graphs( capture_decode_cuda_graph=capture_decode_cuda_graph ) + def _maybe_build_draft_sampler(self): + def _eager(reason): + if self.tp_rank == 0: + logger.info("DFLASH draft greedy head kept eager (reason=%s).", reason) + return None + + if get_tp_group().world_size != 1: + return _eager("tp>1") + if self.block_size <= 1: + return _eager("block_size<=1") + target_model = self._target_worker.model_runner.model + lm_head = getattr(target_model, "lm_head", None) + if lm_head is None or not hasattr(lm_head, "weight"): + return _eager("no target lm_head") + if not torch.is_floating_point(lm_head.weight): + # Quantized lm_head (FP8/INT) would break the static matmul. + return _eager("quantized lm_head") + if not hasattr(lm_head, "shard_indices"): + num_org = int(lm_head.weight.shape[0]) + org_vocab_start = 0 + else: + shard = lm_head.shard_indices + if int(shard.num_added_elements) != 0: + return _eager("added vocab") + num_org = int(shard.num_org_elements) + org_vocab_start = int(shard.org_vocab_start_index) + if self.tp_rank == 0: + logger.info("DFLASH draft greedy head folded into the draft cuda graph.") + return _DflashDraftSampler( + weight=lm_head.weight, + block_size=self.block_size, + num_org=num_org, + org_vocab_start=org_vocab_start, + max_bs=max(self.server_args.cuda_graph_config.decode.bs), + ) + def _init_fused_kv_helper(self) -> None: """Initialize the fused KV materialization helper with pre-stacked weights.""" try: @@ -1480,18 +1555,24 @@ class DFlashWorkerV2(BaseSpecWorker): ) with torch.inference_mode(): - draft_logits_output = self.draft_model_runner.forward( - forward_batch - ).logits_output + draft_out = self.draft_model_runner.forward(forward_batch) + draft_logits_output = draft_out.logits_output - draft_hidden = draft_logits_output.hidden_states - if draft_hidden is None: - raise RuntimeError("DFLASH draft model returned no hidden states.") - draft_hidden = draft_hidden.view(bs, int(self.block_size), -1) - draft_next = self._greedy_sample_from_vocab_parallel_head( - hidden_states=draft_hidden[:, 1:, :].reshape(-1, draft_hidden.shape[-1]), - lm_head=lm_head, - ).view(bs, int(self.block_size) - 1) + if self._draft_sampler is not None and draft_out.can_run_graph: + draft_next = self._draft_sampler.out[ + : bs * (int(self.block_size) - 1) + ].view(bs, int(self.block_size) - 1) + else: + draft_hidden = draft_logits_output.hidden_states + if draft_hidden is None: + raise RuntimeError("DFLASH draft model returned no hidden states.") + draft_hidden = draft_hidden.view(bs, int(self.block_size), -1) + draft_next = self._greedy_sample_from_vocab_parallel_head( + hidden_states=draft_hidden[:, 1:, :].reshape( + -1, draft_hidden.shape[-1] + ), + lm_head=lm_head, + ).view(bs, int(self.block_size) - 1) draft_tokens = self._draft_block_tokens_buf[:bs] draft_tokens[:, 0].copy_(block_ids[:, 0])