diff --git a/python/sglang/srt/environ.py b/python/sglang/srt/environ.py index 2dbc6a04f..30c796851 100644 --- a/python/sglang/srt/environ.py +++ b/python/sglang/srt/environ.py @@ -226,6 +226,16 @@ class InvariantCheckLevel(IntEnum): STRICT = 2 +class DsparkFoldedSampling(IntEnum): + """Sampling support in the graph-folded DSpark draft proposal: OFF = + greedy-only folding, AUTO = on when its buffers fit in free GPU memory, + FORCE = always.""" + + OFF = 0 + AUTO = 1 + FORCE = 2 + + class Envs: # Raise on bare server_args field assignments after resolution; mutation @@ -326,6 +336,7 @@ class Envs: SGLANG_DSPARK_FAST_KERNEL = EnvBool(True) SGLANG_DSPARK_FP32_LM_HEAD = EnvBool(False) SGLANG_DSPARK_FAST_SAMPLING = EnvBool(True) + SGLANG_DSPARK_FOLDED_SAMPLING = EnvInt(DsparkFoldedSampling.AUTO) SGLANG_DSPARK_OPT_MARKOV_W2_BF16 = EnvBool(True) SGLANG_DSPARK_OPT_MARKOV_W2_TP_SHARD = EnvBool(True) SGLANG_DSPARK_ENABLE_MULTI_STREAM = EnvBool(True) diff --git a/python/sglang/srt/speculative/dspark_components/dspark_draft.py b/python/sglang/srt/speculative/dspark_components/dspark_draft.py index acdbd7451..62468e61a 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_draft.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_draft.py @@ -71,90 +71,6 @@ class DraftProposal(msgspec.Struct, frozen=True): folded: bool = False -def greedy_step_sampler(step_logits: torch.Tensor, step_idx: int) -> torch.Tensor: - del step_idx - return torch.argmax(step_logits, dim=-1) - - -class DsparkDraftSampler: - - def __init__(self, *, model, gamma, max_bs, device, confidence_fn=None, out=None): - self.model = model - self.markov_head = model.markov_head - self.gamma = int(gamma) - if out is not None: - assert out.shape == (int(max_bs) * self.gamma,) and out.dtype == torch.int64 - self.out = out - else: - self.out = torch.empty( - (int(max_bs) * self.gamma,), dtype=torch.int64, device=device - ) - self.confidence_fn = confidence_fn - self.confidence_out = ( - torch.empty((int(max_bs), self.gamma), dtype=torch.float32, device=device) - if confidence_fn is not None - else None - ) - - def __call__(self, hidden_states, input_ids): - bs = hidden_states.shape[0] // self.gamma - base_logits, confidence_tap = self.model.compute_base_logits(hidden_states) - base_logits = base_logits.view(bs, self.gamma, -1) - anchor = input_ids.view(bs, self.gamma)[:, 0] - draft_tokens, _ = self.markov_head.sample_block( - base_logits, - first_prev_tokens=anchor, - hidden_states=hidden_states.view(bs, self.gamma, -1), - sampler=greedy_step_sampler, - ) - self.out[: draft_tokens.numel()].copy_(draft_tokens.reshape(-1)) - if self.confidence_out is not None: - confidence = self.confidence_fn( - draft_hidden=hidden_states.view(bs, self.gamma, -1), - anchor_tokens=anchor, - draft_tokens=draft_tokens, - confidence_tap=confidence_tap, - ) - self.confidence_out[:bs].copy_(confidence) - - -def maybe_build_draft_sampler( - *, - draft_model, - gamma: int, - max_bs: int, - device, - tp_rank: int, - confidence_fn=None, - out=None, -) -> Optional[DsparkDraftSampler]: - """Build the graph-folded greedy draft sampler, or return None (with the - reason logged) when the draft model cannot support folding and the - proposal must stay eager.""" - - def _eager(reason): - if tp_rank == 0: - logger.info("DSpark draft greedy proposal kept eager (reason=%s).", reason) - return None - - if gamma <= 0: - return _eager("gamma<=0") - if not hasattr(draft_model, "compute_base_logits"): - return _eager("no compute_base_logits") - if getattr(draft_model, "markov_head", None) is None: - return _eager("no markov head") - if tp_rank == 0: - logger.info("DSpark draft greedy proposal folded into the draft cuda graph.") - return DsparkDraftSampler( - model=draft_model, - gamma=gamma, - max_bs=max_bs, - device=device, - confidence_fn=confidence_fn, - out=out, - ) - - def make_next_draft_input( *, bonus_tokens: torch.Tensor, @@ -277,6 +193,8 @@ class DraftBlockProposer: sampling_info, ) -> DraftProposal: embed_module = target_model.get_input_embeddings() + draft_sampler = self._draft_sampler + all_greedy = sampling_info is None or sampling_info.is_all_greedy fwd = self._run_forward( batch=batch, draft_input=draft_input, @@ -284,30 +202,51 @@ class DraftBlockProposer: bs=bs, device=device, embed_module=embed_module, + draft_sampler=draft_sampler, + sampling_info=sampling_info, ) draft_block_ids = fwd.draft_block_ids - draft_sampler = self._draft_sampler - all_greedy = sampling_info is None or sampling_info.is_all_greedy folded_confidence = None confidence_tap = None folded = False - if draft_sampler is not None and fwd.can_run_graph and all_greedy: + if ( + draft_sampler is not None + and fwd.can_run_graph + and (all_greedy or draft_sampler.folded_sampling) + ): folded = True - if sampling_info is None: - temperatures = torch.ones(bs, dtype=torch.float32, device=device) - else: - temperatures = ( - sampling_info.temperatures.view(-1) - .to(torch.float32) - .clamp_min(1e-5) + if draft_sampler.folded_sampling: + greedy_mask = draft_sampler.greedy_mask[:bs] + temperatures = draft_sampler.temperatures[:bs] + # The sampling accept path needs the markov-corrected block + # logits; greedy accept only compares tokens. + corrected_logits = ( + None + if all_greedy + else draft_sampler.corrected_out[: bs * self.gamma].view( + bs, self.gamma, -1 + ) ) + else: + # Greedy-only folding: the hook argmaxed every row and kept no + # sampling buffers, so derive the params on the fly. + greedy_mask = resolve_greedy_mask( + bs=bs, sampling_info=sampling_info, device=device + ) + if sampling_info is None: + temperatures = torch.ones(bs, dtype=torch.float32, device=device) + else: + temperatures = ( + sampling_info.temperatures.view(-1) + .to(torch.float32) + .clamp_min(1e-5) + ) + corrected_logits = None draft_block = DraftBlockResult( draft_tokens=draft_sampler.out[: bs * self.gamma].view(bs, self.gamma), - corrected_logits=None, - greedy_mask=resolve_greedy_mask( - bs=bs, sampling_info=sampling_info, device=device - ), + corrected_logits=corrected_logits, + greedy_mask=greedy_mask, temperatures=temperatures, ) if draft_sampler.confidence_out is not None: @@ -367,6 +306,8 @@ class DraftBlockProposer: bs: int, device: str, embed_module, + draft_sampler=None, + sampling_info=None, ) -> DraftForwardResult: gamma = self.gamma prefix_lens = batch.seq_lens @@ -411,6 +352,13 @@ class DraftBlockProposer: capture_hidden_mode=CaptureHiddenMode.NULL, ) self._fill_dp_moe_sync_metadata(draft_forward_batch, batch) + graph_runner = self.draft_model_runner.decode_cuda_graph_runner + if ( + draft_sampler is not None + and graph_runner is not None + and graph_runner.can_run_graph(draft_forward_batch) + ): + draft_sampler.stage_sampling_params(bs=bs, sampling_info=sampling_info) with torch.inference_mode(): draft_out = self.draft_model_runner.forward(draft_forward_batch) logits_output = draft_out.logits_output diff --git a/python/sglang/srt/speculative/dspark_components/dspark_draft_sampler.py b/python/sglang/srt/speculative/dspark_components/dspark_draft_sampler.py new file mode 100644 index 000000000..a01c93a7a --- /dev/null +++ b/python/sglang/srt/speculative/dspark_components/dspark_draft_sampler.py @@ -0,0 +1,205 @@ +from __future__ import annotations + +import logging +from typing import Optional + +import torch + +from sglang.kernels.ops.speculative.dspark.dspark_draft_model import ( + SampleStepTokens, +) +from sglang.srt.environ import DsparkFoldedSampling, envs +from sglang.srt.utils import get_available_gpu_memory + +logger = logging.getLogger(__name__) + +# Same free-memory floor init_cuda_graphs requires before draft capture. +_CAPTURE_HEADROOM_GB = 1.0 + + +def greedy_step_sampler(step_logits: torch.Tensor, step_idx: int) -> torch.Tensor: + del step_idx + return torch.argmax(step_logits, dim=-1) + + +class DsparkDraftSampler: + """Draft proposal head folded into the draft graph as a tail hook; with + folded_sampling it also Gumbel-samples non-greedy rows in-graph.""" + + def __init__( + self, + *, + model, + gamma, + max_bs, + device, + confidence_fn=None, + out=None, + folded_sampling: bool = True, + ): + self.model = model + self.markov_head = model.markov_head + self.gamma = int(gamma) + max_bs = int(max_bs) + if out is not None: + assert out.shape == (max_bs * self.gamma,) and out.dtype == torch.int64 + self.out = out + else: + self.out = torch.empty( + (max_bs * self.gamma,), dtype=torch.int64, device=device + ) + self.confidence_fn = confidence_fn + self.confidence_out = ( + torch.empty((max_bs, self.gamma), dtype=torch.float32, device=device) + if confidence_fn is not None + else None + ) + self.folded_sampling = folded_sampling + self.temperatures = None + self.greedy_mask = None + self.exp_noise = None + self.corrected_out = None + if folded_sampling: + vocab = int(model.lm_head.org_vocab_size) + self.temperatures = torch.ones( + (max_bs,), dtype=torch.float32, device=device + ) + self.greedy_mask = torch.ones((max_bs,), dtype=torch.bool, device=device) + self.exp_noise = torch.empty( + (max_bs, vocab), dtype=torch.float32, device=device + ) + self.corrected_out = torch.empty( + (max_bs * self.gamma, vocab), + dtype=model.lm_head.weight.dtype, + device=device, + ) + + def stage_sampling_params(self, *, bs: int, sampling_info) -> None: + """Host-side refresh of the static sampling params; must run before + the draft graph replay that consumes them.""" + if not self.folded_sampling: + return + if sampling_info is None: + self.temperatures[:bs].fill_(1.0) + self.greedy_mask[:bs].fill_(True) + return + torch.clamp( + sampling_info.temperatures.view(-1)[:bs].to(torch.float32), + min=1e-5, + out=self.temperatures[:bs], + ) + self.greedy_mask[:bs].copy_((sampling_info.top_ks <= 1).view(-1)[:bs]) + + def __call__(self, hidden_states, input_ids): + bs = hidden_states.shape[0] // self.gamma + base_logits, confidence_tap = self.model.compute_base_logits(hidden_states) + base_logits = base_logits.view(bs, self.gamma, -1) + anchor = input_ids.view(bs, self.gamma)[:, 0] + + if self.folded_sampling: + + def sampler(step_logits: torch.Tensor, step_idx: int) -> torch.Tensor: + del step_idx + # In-graph philox noise: each replay advances the generator + # and redraws. + noise = self.exp_noise[:bs].exponential_() + return SampleStepTokens.execute( + step_logits=step_logits, + temperatures=self.temperatures[:bs], + greedy_mask=self.greedy_mask[:bs], + exp_noise=noise, + ) + + else: + sampler = greedy_step_sampler + + draft_tokens, corrected_logits = self.markov_head.sample_block( + base_logits, + first_prev_tokens=anchor, + hidden_states=hidden_states.view(bs, self.gamma, -1), + sampler=sampler, + ) + self.out[: draft_tokens.numel()].copy_(draft_tokens.reshape(-1)) + if self.folded_sampling: + self.corrected_out[: bs * self.gamma].copy_( + corrected_logits.reshape(bs * self.gamma, -1) + ) + if self.confidence_out is not None: + confidence = self.confidence_fn( + draft_hidden=hidden_states.view(bs, self.gamma, -1), + anchor_tokens=anchor, + draft_tokens=draft_tokens, + confidence_tap=confidence_tap, + ) + self.confidence_out[:bs].copy_(confidence) + + +def _resolve_folded_sampling(*, model, gamma, max_bs, device, tp_rank) -> bool: + """The sampling buffers are baked into the captured draft graph, so AUTO + must decide before capture from a free-memory probe.""" + mode = envs.SGLANG_DSPARK_FOLDED_SAMPLING.get() + if mode == DsparkFoldedSampling.OFF: + return False + if mode == DsparkFoldedSampling.FORCE: + return True + vocab = int(model.lm_head.org_vocab_size) + noise_bytes = max_bs * vocab * 4 + logits_bytes = max_bs * gamma * vocab * model.lm_head.weight.dtype.itemsize + need_gb = (noise_bytes + logits_bytes) / (1 << 30) + available_gb = get_available_gpu_memory(device, torch.cuda.current_device()) + if available_gb - need_gb >= _CAPTURE_HEADROOM_GB: + return True + if tp_rank == 0: + logger.warning( + "DSpark folded sampling disabled: its static buffers need %.2f GB " + "but only %.2f GB GPU memory is free; sampling batches will take " + "the eager proposal path. Set SGLANG_DSPARK_FOLDED_SAMPLING=%d " + "to force.", + need_gb, + available_gb, + int(DsparkFoldedSampling.FORCE), + ) + return False + + +def maybe_build_draft_sampler( + *, + draft_model, + gamma: int, + max_bs: int, + device, + tp_rank: int, + confidence_fn=None, + out=None, +) -> Optional[DsparkDraftSampler]: + """Build the graph-folded draft sampler, or None (reason logged) when the + proposal must stay eager.""" + + def _eager(reason): + if tp_rank == 0: + logger.info("DSpark draft proposal kept eager (reason=%s).", reason) + return None + + if gamma <= 0: + return _eager("gamma<=0") + if not hasattr(draft_model, "compute_base_logits"): + return _eager("no compute_base_logits") + if getattr(draft_model, "markov_head", None) is None: + return _eager("no markov head") + folded_sampling = _resolve_folded_sampling( + model=draft_model, gamma=gamma, max_bs=max_bs, device=device, tp_rank=tp_rank + ) + if tp_rank == 0: + logger.info( + "DSpark draft proposal (%s) folded into the draft cuda graph.", + "greedy + sampling" if folded_sampling else "greedy only", + ) + return DsparkDraftSampler( + model=draft_model, + gamma=gamma, + max_bs=max_bs, + device=device, + confidence_fn=confidence_fn, + out=out, + folded_sampling=folded_sampling, + ) diff --git a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py index a72ea6360..74663aebf 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py @@ -33,6 +33,8 @@ from sglang.srt.speculative.dspark_components.dspark_config import ( from sglang.srt.speculative.dspark_components.dspark_draft import ( DraftBlockProposer, make_next_draft_input, +) +from sglang.srt.speculative.dspark_components.dspark_draft_sampler import ( maybe_build_draft_sampler, ) from sglang.srt.speculative.dspark_components.dspark_kv_inject import ( @@ -592,6 +594,10 @@ class DSparkWorkerV2(BaseSpecWorker): fold_eligible = ( self._verify_executor.verify_epilogue is not None and proposal.folded + # The epilogue's in-graph accept is greedy (accept_greedy_triton); + # sampling batches must take the eager accept path even when the + # draft proposal itself folded. + and (sampling_info is None or sampling_info.is_all_greedy) and verify_logits_adjustments_are_noop(sampling_info) and self._simulate_acc_len <= 0 and not batch.has_grammar