[Spec] Deduplicate spec-v2 worker lifecycle boilerplate into BaseSpecWorker (#31008)

This commit is contained in:
Liangsheng Yin
2026-07-13 13:48:40 -05:00
committed by GitHub
parent c0f1f7e062
commit 48fff1f2bd
9 changed files with 56 additions and 135 deletions
@@ -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
@@ -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
@@ -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 (
@@ -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
@@ -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:
@@ -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
@@ -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.
+13 -1
View File
@@ -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()
@@ -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