[Spec] Deduplicate spec-v2 worker lifecycle boilerplate into BaseSpecWorker (#31008)
This commit is contained in:
@@ -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.
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user