Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
11 changes: 7 additions & 4 deletions python/sglang/srt/arg_groups/speculative_hook.py
Original file line number Diff line number Diff line change
Expand Up @@ -256,10 +256,13 @@ def _handle_frozen_kv_mtp(server_args: "ServerArgs") -> None:
"Max running requests is reset to 48 for speculative decoding. You can override this by explicitly setting --max-running-requests."
)

server_args.disable_overlap_schedule = True
logger.warning(
"Overlap scheduler is disabled when using Frozen-KV MTP speculative decoding (spec v2 is not supported yet)."
)
# SGLANG_ENABLE_SPEC_V2=False selects the non-overlap (synchronous) spec v2
# path instead of the overlap-scheduled one; both run the V2 worker.
if (
not envs.SGLANG_ENABLE_SPEC_V2.get()
and not server_args.disable_overlap_schedule
):
server_args.disable_overlap_schedule = True

if server_args.enable_mixed_chunk:
server_args.enable_mixed_chunk = False
Expand Down
8 changes: 4 additions & 4 deletions python/sglang/srt/managers/scheduler.py
Original file line number Diff line number Diff line change
Expand Up @@ -1163,7 +1163,7 @@ def init_overlap(self):
self.device_module = torch.get_device_module(self.device)

# FutureMap is always-on: input_ids relay used in both modes.
# Workers not on BaseSpecWorker (e.g. FrozenKVMTPWorker) lack the
# Workers not on BaseSpecWorker (e.g. NGRAM / DFLASH) lack the
# override; fall back to target-only so the helper still produces a
# safe decision (no accidental opt-out for unaudited shapes).
if self.draft_worker is not None:
Expand Down Expand Up @@ -3120,9 +3120,9 @@ def run_batch(
)
batch.input_ids = None
else:
# Spec_v1 (NGRAM / DFLASH / FROZEN_KV_MTP, non-overlap):
# worker shape doesn't match req_pool_indices; relay is
# unused (worker rebuilds input_ids inside verify).
# Spec_v1 (NGRAM / DFLASH, non-overlap): worker shape
# doesn't match req_pool_indices; relay is unused (worker
# rebuilds input_ids inside verify).
batch.input_ids = batch_result.next_token_ids.to(torch.int64)
self.update_cache_from_scheduler(batch, batch_result)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -34,7 +34,7 @@
)

if TYPE_CHECKING:
from sglang.srt.speculative.frozen_kv_mtp_worker import FrozenKVMTPWorker
from sglang.srt.speculative.frozen_kv_mtp_worker_v2 import FrozenKVMTPDraftWorker


@dataclass
Expand All @@ -47,7 +47,7 @@ class FrozenKVMTPInputBuffers(ForwardInputBuffers):
topk_p: torch.Tensor
topk_index: torch.Tensor
hidden_states: torch.Tensor
# Consumed by the captured seed iter; see `FrozenKVMTPWorker.draft_forward`.
# Consumed by the captured seed iter; see `FrozenKVMTPDraftWorker.draft_forward`.
bonus_tokens: torch.Tensor
global_num_tokens_gpu: Optional[torch.Tensor]
global_num_tokens_for_logprob_gpu: Optional[torch.Tensor]
Expand All @@ -56,7 +56,7 @@ class FrozenKVMTPInputBuffers(ForwardInputBuffers):
class FrozenKVMTPCudaGraphRunner:
"""CUDA graph runner for the Frozen-KV MTP recurrent draft-loop step."""

def __init__(self, frozen_kv_mtp_worker: FrozenKVMTPWorker):
def __init__(self, frozen_kv_mtp_worker: FrozenKVMTPDraftWorker):
self.frozen_kv_mtp_worker = frozen_kv_mtp_worker
self.model_runner = model_runner = frozen_kv_mtp_worker.draft_model_runner
self.graphs = {}
Expand Down
26 changes: 1 addition & 25 deletions python/sglang/srt/speculative/frozen_kv_mtp_info.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,15 +13,14 @@
# ==============================================================================
from __future__ import annotations

from dataclasses import dataclass, fields
from dataclasses import dataclass
from typing import Dict

from sglang.srt.mem_cache.memory_pool import KVCache
from sglang.srt.speculative.eagle_info import (
EagleDraftExtendInput,
EagleDraftInput,
EagleVerifyInput,
EagleVerifyOutput,
)
from sglang.srt.speculative.spec_info import SpecInput, SpecInputType

Expand Down Expand Up @@ -68,26 +67,3 @@ class FrozenKVMTPVerifyInput(EagleVerifyInput):

def __post_init__(self):
SpecInput.__init__(self, SpecInputType.FROZEN_KV_MTP_VERIFY)

def verify(self, *args, **kwargs) -> EagleVerifyOutput:
output = super().verify(*args, **kwargs)
output.draft_extend_input = _to_frozen_kv_mtp_draft_extend_input(
output.draft_extend_input
)
return output


FrozenKVMTPVerifyOutput = EagleVerifyOutput


def _to_frozen_kv_mtp_draft_extend_input(
draft_extend_input: EagleDraftExtendInput,
) -> FrozenKVMTPDraftExtendInput:
if isinstance(draft_extend_input, FrozenKVMTPDraftExtendInput):
return draft_extend_input
return FrozenKVMTPDraftExtendInput(
**{
field.name: getattr(draft_extend_input, field.name)
for field in fields(EagleDraftExtendInput)
}
)
29 changes: 2 additions & 27 deletions python/sglang/srt/speculative/frozen_kv_mtp_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,19 +14,13 @@
from __future__ import annotations

from contextlib import contextmanager
from typing import TYPE_CHECKING, Tuple
from typing import TYPE_CHECKING

import torch

from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
from sglang.srt.speculative.frozen_kv_mtp_info import (
FrozenKVMTPContext,
FrozenKVMTPDraftExtendInput,
FrozenKVMTPDraftInput,
)
from sglang.srt.speculative.spec_utils import fast_topk
from sglang.srt.speculative.frozen_kv_mtp_info import FrozenKVMTPContext

if TYPE_CHECKING:
from sglang.srt.layers.attention.base_attn_backend import AttentionBackend
Expand Down Expand Up @@ -159,22 +153,3 @@ def select_last_extend_hidden(
lens = torch.tensor(batch.extend_lens, device=hidden_states.device)
last_indices = torch.cumsum(lens, dim=0) - 1
return hidden_states[last_indices.to(torch.long)]


def select_last_verified_seed(
draft_input: FrozenKVMTPDraftExtendInput,
) -> Tuple[torch.Tensor, torch.Tensor]:
counts = draft_input.num_accept_tokens.to(torch.long)
last_indices = torch.cumsum(counts, dim=0) - 1
return (
draft_input.input_ids[last_indices],
draft_input.hidden_states[last_indices],
)


def capture_for_decode(
logits_output: LogitsProcessorOutput, draft_input: FrozenKVMTPDraftInput, topk: int
) -> None:
probs = torch.softmax(logits_output.next_token_logits, dim=-1)
draft_input.topk_p, draft_input.topk_index = fast_topk(probs, topk, dim=-1)
draft_input.hidden_states = logits_output.hidden_states
Loading
Loading