diff --git a/python/sglang/srt/speculative/base_spec_worker.py b/python/sglang/srt/speculative/base_spec_worker.py index 95d28026b..4268fea18 100644 --- a/python/sglang/srt/speculative/base_spec_worker.py +++ b/python/sglang/srt/speculative/base_spec_worker.py @@ -1,7 +1,7 @@ from __future__ import annotations from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING, Any, Optional import torch @@ -291,14 +291,14 @@ class EagleDraftWorkerBase(ABC): class BaseSpecWorker(ABC): @property - @abstractmethod def target_worker(self) -> TpModelWorker: - pass + return self._target_worker @property - @abstractmethod - def draft_worker(self) -> EagleDraftWorkerBase: - pass + def draft_worker(self) -> Optional[EagleDraftWorkerBase | TpModelWorker]: + # dflash / dspark drive the draft model through a plain TpModelWorker; + # ngram has no draft worker at all (returns None via its override). + return self._draft_worker @property def war_fastpath_runner(self): @@ -314,19 +314,34 @@ class BaseSpecWorker(ABC): Default returns target only; subclasses extend with draft backends.""" return (self.target_worker.model_runner.attn_backend,) - @abstractmethod def clear_cache_pool(self): - # TODO: move this abstract method to BaseTpWorker and call through self.model_runner + """Default no-op: the allocator and kv cache pool are shared with the + target worker and cleared by the scheduler.""" + # TODO: move this method to BaseTpWorker and call through self.model_runner pass - def alloc_memory_pool(self, **kwargs): - pass + def alloc_memory_pool( + self, + memory_pool_config=None, + req_to_token_pool=None, + token_to_kv_pool_allocator=None, + ): + if self.draft_worker is not None: + self.draft_worker.alloc_memory_pool( + memory_pool_config=memory_pool_config, + req_to_token_pool=req_to_token_pool, + token_to_kv_pool_allocator=token_to_kv_pool_allocator, + ) + self.req_to_token_pool = req_to_token_pool + self.token_to_kv_pool_allocator = token_to_kv_pool_allocator def init_attention_backends(self): - pass + if self.draft_worker is not None: + self.draft_worker.init_attention_backends() def init_cuda_graphs(self): - pass + if self.draft_worker is not None: + self.draft_worker.init_cuda_graphs() def on_verify_complete_cpu( self, num_correct_drafts_per_req: list[int], batch_size: int = 0 diff --git a/python/sglang/srt/speculative/dflash_worker_v2.py b/python/sglang/srt/speculative/dflash_worker_v2.py index e59c4d8d6..d9c82b360 100644 --- a/python/sglang/srt/speculative/dflash_worker_v2.py +++ b/python/sglang/srt/speculative/dflash_worker_v2.py @@ -240,10 +240,6 @@ class DFlashWorkerV2(BaseSpecWorker): self._out_tokens_bufs: List[torch.Tensor] = [] self._new_seq_lens_bufs: List[torch.Tensor] = [] - @property - def target_worker(self) -> TpModelWorker: - return self._target_worker - @property def draft_worker(self): # DFLASH drives the draft model through a plain TpModelWorker: the 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 ee525f02a..6e3f8da86 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py @@ -270,14 +270,6 @@ class DSparkWorkerV2(BaseSpecWorker): def carries_confidence(self) -> bool: return self._verify_planner.carries_confidence - @property - def target_worker(self) -> TpModelWorker: - return self._target_worker - - @property - def draft_worker(self): - return self._draft_worker - @property def spec_v2_attn_backends(self) -> tuple: return ( diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 5c8478f16..f0fde3aaa 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -1,7 +1,7 @@ import contextlib import logging import time -from typing import List, Optional, Tuple +from typing import List, Optional import torch @@ -82,6 +82,7 @@ from sglang.srt.speculative.spec_utils import ( draft_tp_context, fast_sample, generate_token_bitmask, + get_plan_stream, load_token_map, move_accept_tokens_to_target_kvcache, record_stream_each, @@ -122,17 +123,6 @@ _is_xpu = is_xpu() logger = logging.getLogger(__name__) -def _get_plan_stream( - device: str, -) -> Tuple[any, contextlib.AbstractContextManager]: - if envs.SGLANG_ENABLE_OVERLAP_PLAN_STREAM.get(): - plan_stream = torch.get_device_module(device).Stream() - plan_stream_ctx = torch.get_device_module(device).stream(plan_stream) - return plan_stream, plan_stream_ctx - else: - return None, contextlib.nullcontext() - - class EagleDraftWorker(EagleDraftWorkerBase): def __init__( self, @@ -204,7 +194,7 @@ class EagleDraftWorker(EagleDraftWorkerBase): ) self.tree_mask_mode = default_tree_mask_mode() - self.plan_stream, self.plan_stream_ctx = _get_plan_stream(self.device) + self.plan_stream, self.plan_stream_ctx = get_plan_stream(self.device) def alloc_memory_pool( self, @@ -1107,7 +1097,7 @@ class EAGLEWorkerV2(BaseSpecWorker): ) self.extend_lens = torch.empty((), dtype=torch.int64, device=self.device) - self.plan_stream, self.plan_stream_ctx = _get_plan_stream(self.device) + self.plan_stream, self.plan_stream_ctx = get_plan_stream(self.device) @property def war_fastpath_runner(self): @@ -1126,23 +1116,8 @@ class EAGLEWorkerV2(BaseSpecWorker): or self._draft_worker.draft_runner.attn_backend, ) - def alloc_memory_pool( - self, - memory_pool_config=None, - req_to_token_pool=None, - token_to_kv_pool_allocator=None, - ): - self._draft_worker.alloc_memory_pool( - memory_pool_config, req_to_token_pool, token_to_kv_pool_allocator - ) - self.req_to_token_pool = req_to_token_pool - self.token_to_kv_pool_allocator = token_to_kv_pool_allocator - - def init_attention_backends(self): - self._draft_worker.init_attention_backends() - def init_cuda_graphs(self): - self._draft_worker.init_cuda_graphs() + super().init_cuda_graphs() # Build adaptive runtime states after target and draft backends exist. if self.adaptive_controller is not None: with ( @@ -1172,18 +1147,6 @@ class EAGLEWorkerV2(BaseSpecWorker): ), ) - @property - def target_worker(self): - return self._target_worker - - @property - def draft_worker(self): - return self._draft_worker - - def clear_cache_pool(self): - # allocator and kv cache pool are shared with target worker, which are cleared in scheduler - pass - def forward_batch_generation(self, batch: ScheduleBatch, on_publish=None): if batch.forward_mode.is_extend() or batch.is_extend_in_batch: # Target prefill diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py b/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py index 6c74f30a1..c9ec82ff5 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_worker_v2.py @@ -46,7 +46,7 @@ from sglang.srt.speculative.eagle_utils import ( build_tree_kernel_efficient, organize_draft_results, ) -from sglang.srt.speculative.eagle_worker_v2 import EAGLEWorkerV2, _get_plan_stream +from sglang.srt.speculative.eagle_worker_v2 import EAGLEWorkerV2 from sglang.srt.speculative.frozen_kv_mtp_info import ( FrozenKVMTPContext, FrozenKVMTPDraftInput, @@ -64,6 +64,7 @@ from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_utils import ( draft_tp_context, fast_topk, + get_plan_stream, select_top_k_tokens, spec_stage_span, ) @@ -705,7 +706,7 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2): ) self.extend_lens = torch.empty((), dtype=torch.int64, device=self.device) - self.plan_stream, self.plan_stream_ctx = _get_plan_stream(self.device) + self.plan_stream, self.plan_stream_ctx = get_plan_stream(self.device) @property def spec_v2_attn_backends(self) -> tuple: diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index f98298721..f0dea5595 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -12,9 +12,8 @@ # limitations under the License. # ============================================================================== -import contextlib import logging -from typing import TYPE_CHECKING, List, Optional, Tuple +from typing import TYPE_CHECKING, List, Optional import torch @@ -63,6 +62,7 @@ from sglang.srt.speculative.multi_layer_eagle_utils import rotate_input_ids from sglang.srt.speculative.spec_info import SpeculativeAlgorithm from sglang.srt.speculative.spec_utils import ( draft_tp_context, + get_plan_stream, record_stream_each, record_stream_for_v2_verify, sample_draft_proposal, @@ -87,17 +87,6 @@ if TYPE_CHECKING: logger = logging.getLogger(__name__) -def _get_plan_stream( - device: str, -) -> Tuple[any, contextlib.AbstractContextManager]: - if envs.SGLANG_ENABLE_OVERLAP_PLAN_STREAM.get(): - plan_stream = torch.get_device_module(device).Stream() - plan_stream_ctx = torch.get_device_module(device).stream(plan_stream) - return plan_stream, plan_stream_ctx - else: - return None, contextlib.nullcontext() - - class MultiLayerEagleDraftWorker(EagleDraftWorkerBase): def __init__( self, @@ -171,7 +160,7 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase): draft_tp_context if server_args.enable_dp_attention else empty_context ) self.tree_mask_mode = default_tree_mask_mode() - self.plan_stream, self.plan_stream_ctx = _get_plan_stream(self.device) + self.plan_stream, self.plan_stream_ctx = get_plan_stream(self.device) def alloc_memory_pool( self, @@ -707,33 +696,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker): ) self.extend_lens = torch.empty((), dtype=torch.int64, device=self.device) - self.plan_stream, self.plan_stream_ctx = _get_plan_stream(self.device) - - def alloc_memory_pool( - self, - memory_pool_config=None, - req_to_token_pool=None, - token_to_kv_pool_allocator=None, - ): - self._draft_worker.alloc_memory_pool( - memory_pool_config, req_to_token_pool, token_to_kv_pool_allocator - ) - self.req_to_token_pool = req_to_token_pool - self.token_to_kv_pool_allocator = token_to_kv_pool_allocator - - def init_attention_backends(self): - self._draft_worker.init_attention_backends() - - def init_cuda_graphs(self): - self._draft_worker.init_cuda_graphs() - - @property - def target_worker(self): - return self._target_worker - - @property - def draft_worker(self): - return self._draft_worker + self.plan_stream, self.plan_stream_ctx = get_plan_stream(self.device) @property def spec_v2_attn_backends(self) -> tuple: @@ -748,10 +711,6 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker): ), ) - def clear_cache_pool(self): - # allocator and kv cache pool are shared with target worker, which are cleared in scheduler - pass - def forward_batch_generation(self, batch: ScheduleBatch, on_publish=None): if batch.forward_mode.is_extend() or batch.is_extend_in_batch: # Target prefill diff --git a/python/sglang/srt/speculative/ngram_worker.py b/python/sglang/srt/speculative/ngram_worker.py index ed1d657fb..6a875b397 100644 --- a/python/sglang/srt/speculative/ngram_worker.py +++ b/python/sglang/srt/speculative/ngram_worker.py @@ -109,10 +109,6 @@ class NGRAMWorker(BaseSpecWorker): loaded, ) - @property - def target_worker(self) -> TpModelWorker: - return self._target_worker - @property def draft_worker(self) -> Optional[EagleDraftWorkerBase]: # NGRAM has no draft model; drafts come from the CPU-side corpus. diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index 70a47954d..d170dc634 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -1,10 +1,11 @@ from __future__ import annotations +import contextlib import logging import os import time from contextlib import contextmanager -from typing import TYPE_CHECKING, List, Optional +from typing import TYPE_CHECKING, Any, List, Optional, Tuple import torch from huggingface_hub import snapshot_download @@ -717,3 +718,14 @@ def spec_prepare_for_decode(batch: ScheduleBatch) -> None: from sglang.srt.speculative.eagle_utils import eagle_prepare_for_decode eagle_prepare_for_decode(batch) + + +def get_plan_stream( + device: str, +) -> Tuple[Any, contextlib.AbstractContextManager]: + if envs.SGLANG_ENABLE_OVERLAP_PLAN_STREAM.get(): + plan_stream = torch.get_device_module(device).Stream() + plan_stream_ctx = torch.get_device_module(device).stream(plan_stream) + return plan_stream, plan_stream_ctx + else: + return None, contextlib.nullcontext() diff --git a/python/sglang/srt/speculative/standalone_worker_v2.py b/python/sglang/srt/speculative/standalone_worker_v2.py index 10dc6f7af..7518da127 100644 --- a/python/sglang/srt/speculative/standalone_worker_v2.py +++ b/python/sglang/srt/speculative/standalone_worker_v2.py @@ -1,10 +1,8 @@ -import contextlib import logging -from typing import Optional, Tuple +from typing import Optional import torch -from sglang.srt.environ import envs from sglang.srt.layers.moe.utils import speculative_moe_backend_context from sglang.srt.managers.tp_worker import TpModelWorker from sglang.srt.server_args import ServerArgs @@ -14,7 +12,7 @@ from sglang.srt.speculative.adaptive_runtime_state import ( from sglang.srt.speculative.eagle_utils import default_tree_mask_mode from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker, EAGLEWorkerV2 from sglang.srt.speculative.spec_info import SpeculativeAlgorithm -from sglang.srt.speculative.spec_utils import draft_tp_context +from sglang.srt.speculative.spec_utils import draft_tp_context, get_plan_stream from sglang.srt.utils import empty_context, get_bool_env_var, is_cuda if is_cuda(): @@ -24,17 +22,6 @@ logger = logging.getLogger(__name__) SGLANG_RETURN_ORIGINAL_LOGPROB = get_bool_env_var("SGLANG_RETURN_ORIGINAL_LOGPROB") -def _get_plan_stream( - device: str, -) -> Tuple[any, contextlib.AbstractContextManager]: - if envs.SGLANG_ENABLE_OVERLAP_PLAN_STREAM.get(): - plan_stream = torch.get_device_module(device).Stream() - plan_stream_ctx = torch.get_device_module(device).stream(plan_stream) - return plan_stream, plan_stream_ctx - else: - return None, contextlib.nullcontext() - - class StandaloneDraftWorker(EagleDraftWorker): """Custom EagleDraftWorker that doesn't share embeddings/lm_head with target model.""" @@ -103,7 +90,7 @@ class StandaloneDraftWorker(EagleDraftWorker): draft_tp_context if server_args.enable_dp_attention else empty_context ) self.tree_mask_mode = default_tree_mask_mode() - self.plan_stream, self.plan_stream_ctx = _get_plan_stream(self.device) + self.plan_stream, self.plan_stream_ctx = get_plan_stream(self.device) # draft_forward reads this (set in EagleDraftWorker.__init__, skipped here). self.index_share_for_mtp_iteration = ( getattr( @@ -210,7 +197,7 @@ class StandaloneWorkerV2(EAGLEWorkerV2): ) self.extend_lens = torch.empty((), dtype=torch.int64, device=self.device) - self.plan_stream, self.plan_stream_ctx = _get_plan_stream(self.device) + self.plan_stream, self.plan_stream_ctx = get_plan_stream(self.device) # TODO: Adaptive speculative self.adaptive_controller: Optional[AdaptiveController] = None