[Spec] Deduplicate spec-v2 worker lifecycle boilerplate into BaseSpecWorker (#31008)
This commit is contained in:
@@ -1,7 +1,7 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import TYPE_CHECKING, Any
|
from typing import TYPE_CHECKING, Any, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -291,14 +291,14 @@ class EagleDraftWorkerBase(ABC):
|
|||||||
|
|
||||||
class BaseSpecWorker(ABC):
|
class BaseSpecWorker(ABC):
|
||||||
@property
|
@property
|
||||||
@abstractmethod
|
|
||||||
def target_worker(self) -> TpModelWorker:
|
def target_worker(self) -> TpModelWorker:
|
||||||
pass
|
return self._target_worker
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@abstractmethod
|
def draft_worker(self) -> Optional[EagleDraftWorkerBase | TpModelWorker]:
|
||||||
def draft_worker(self) -> EagleDraftWorkerBase:
|
# dflash / dspark drive the draft model through a plain TpModelWorker;
|
||||||
pass
|
# ngram has no draft worker at all (returns None via its override).
|
||||||
|
return self._draft_worker
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def war_fastpath_runner(self):
|
def war_fastpath_runner(self):
|
||||||
@@ -314,19 +314,34 @@ class BaseSpecWorker(ABC):
|
|||||||
Default returns target only; subclasses extend with draft backends."""
|
Default returns target only; subclasses extend with draft backends."""
|
||||||
return (self.target_worker.model_runner.attn_backend,)
|
return (self.target_worker.model_runner.attn_backend,)
|
||||||
|
|
||||||
@abstractmethod
|
|
||||||
def clear_cache_pool(self):
|
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
|
pass
|
||||||
|
|
||||||
def alloc_memory_pool(self, **kwargs):
|
def alloc_memory_pool(
|
||||||
pass
|
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):
|
def init_attention_backends(self):
|
||||||
pass
|
if self.draft_worker is not None:
|
||||||
|
self.draft_worker.init_attention_backends()
|
||||||
|
|
||||||
def init_cuda_graphs(self):
|
def init_cuda_graphs(self):
|
||||||
pass
|
if self.draft_worker is not None:
|
||||||
|
self.draft_worker.init_cuda_graphs()
|
||||||
|
|
||||||
def on_verify_complete_cpu(
|
def on_verify_complete_cpu(
|
||||||
self, num_correct_drafts_per_req: list[int], batch_size: int = 0
|
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._out_tokens_bufs: List[torch.Tensor] = []
|
||||||
self._new_seq_lens_bufs: List[torch.Tensor] = []
|
self._new_seq_lens_bufs: List[torch.Tensor] = []
|
||||||
|
|
||||||
@property
|
|
||||||
def target_worker(self) -> TpModelWorker:
|
|
||||||
return self._target_worker
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def draft_worker(self):
|
def draft_worker(self):
|
||||||
# DFLASH drives the draft model through a plain TpModelWorker: the
|
# DFLASH drives the draft model through a plain TpModelWorker: the
|
||||||
|
|||||||
@@ -270,14 +270,6 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
def carries_confidence(self) -> bool:
|
def carries_confidence(self) -> bool:
|
||||||
return self._verify_planner.carries_confidence
|
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
|
@property
|
||||||
def spec_v2_attn_backends(self) -> tuple:
|
def spec_v2_attn_backends(self) -> tuple:
|
||||||
return (
|
return (
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||
import contextlib
|
import contextlib
|
||||||
import logging
|
import logging
|
||||||
import time
|
import time
|
||||||
from typing import List, Optional, Tuple
|
from typing import List, Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
@@ -82,6 +82,7 @@ from sglang.srt.speculative.spec_utils import (
|
|||||||
draft_tp_context,
|
draft_tp_context,
|
||||||
fast_sample,
|
fast_sample,
|
||||||
generate_token_bitmask,
|
generate_token_bitmask,
|
||||||
|
get_plan_stream,
|
||||||
load_token_map,
|
load_token_map,
|
||||||
move_accept_tokens_to_target_kvcache,
|
move_accept_tokens_to_target_kvcache,
|
||||||
record_stream_each,
|
record_stream_each,
|
||||||
@@ -122,17 +123,6 @@ _is_xpu = is_xpu()
|
|||||||
logger = logging.getLogger(__name__)
|
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):
|
class EagleDraftWorker(EagleDraftWorkerBase):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -204,7 +194,7 @@ class EagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
)
|
)
|
||||||
self.tree_mask_mode = default_tree_mask_mode()
|
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(
|
def alloc_memory_pool(
|
||||||
self,
|
self,
|
||||||
@@ -1107,7 +1097,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
)
|
)
|
||||||
self.extend_lens = torch.empty((), dtype=torch.int64, device=self.device)
|
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
|
@property
|
||||||
def war_fastpath_runner(self):
|
def war_fastpath_runner(self):
|
||||||
@@ -1126,23 +1116,8 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
or self._draft_worker.draft_runner.attn_backend,
|
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):
|
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.
|
# Build adaptive runtime states after target and draft backends exist.
|
||||||
if self.adaptive_controller is not None:
|
if self.adaptive_controller is not None:
|
||||||
with (
|
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):
|
def forward_batch_generation(self, batch: ScheduleBatch, on_publish=None):
|
||||||
if batch.forward_mode.is_extend() or batch.is_extend_in_batch:
|
if batch.forward_mode.is_extend() or batch.is_extend_in_batch:
|
||||||
# Target prefill
|
# Target prefill
|
||||||
|
|||||||
@@ -46,7 +46,7 @@ from sglang.srt.speculative.eagle_utils import (
|
|||||||
build_tree_kernel_efficient,
|
build_tree_kernel_efficient,
|
||||||
organize_draft_results,
|
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 (
|
from sglang.srt.speculative.frozen_kv_mtp_info import (
|
||||||
FrozenKVMTPContext,
|
FrozenKVMTPContext,
|
||||||
FrozenKVMTPDraftInput,
|
FrozenKVMTPDraftInput,
|
||||||
@@ -64,6 +64,7 @@ from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
|||||||
from sglang.srt.speculative.spec_utils import (
|
from sglang.srt.speculative.spec_utils import (
|
||||||
draft_tp_context,
|
draft_tp_context,
|
||||||
fast_topk,
|
fast_topk,
|
||||||
|
get_plan_stream,
|
||||||
select_top_k_tokens,
|
select_top_k_tokens,
|
||||||
spec_stage_span,
|
spec_stage_span,
|
||||||
)
|
)
|
||||||
@@ -705,7 +706,7 @@ class FrozenKVMTPWorkerV2(EAGLEWorkerV2):
|
|||||||
)
|
)
|
||||||
self.extend_lens = torch.empty((), dtype=torch.int64, device=self.device)
|
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
|
@property
|
||||||
def spec_v2_attn_backends(self) -> tuple:
|
def spec_v2_attn_backends(self) -> tuple:
|
||||||
|
|||||||
@@ -12,9 +12,8 @@
|
|||||||
# limitations under the License.
|
# limitations under the License.
|
||||||
# ==============================================================================
|
# ==============================================================================
|
||||||
|
|
||||||
import contextlib
|
|
||||||
import logging
|
import logging
|
||||||
from typing import TYPE_CHECKING, List, Optional, Tuple
|
from typing import TYPE_CHECKING, List, Optional
|
||||||
|
|
||||||
import torch
|
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_info import SpeculativeAlgorithm
|
||||||
from sglang.srt.speculative.spec_utils import (
|
from sglang.srt.speculative.spec_utils import (
|
||||||
draft_tp_context,
|
draft_tp_context,
|
||||||
|
get_plan_stream,
|
||||||
record_stream_each,
|
record_stream_each,
|
||||||
record_stream_for_v2_verify,
|
record_stream_for_v2_verify,
|
||||||
sample_draft_proposal,
|
sample_draft_proposal,
|
||||||
@@ -87,17 +87,6 @@ if TYPE_CHECKING:
|
|||||||
logger = logging.getLogger(__name__)
|
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):
|
class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
@@ -171,7 +160,7 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
|
|||||||
draft_tp_context if server_args.enable_dp_attention else empty_context
|
draft_tp_context if server_args.enable_dp_attention else empty_context
|
||||||
)
|
)
|
||||||
self.tree_mask_mode = default_tree_mask_mode()
|
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(
|
def alloc_memory_pool(
|
||||||
self,
|
self,
|
||||||
@@ -707,33 +696,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
)
|
)
|
||||||
self.extend_lens = torch.empty((), dtype=torch.int64, device=self.device)
|
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)
|
||||||
|
|
||||||
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
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def spec_v2_attn_backends(self) -> tuple:
|
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):
|
def forward_batch_generation(self, batch: ScheduleBatch, on_publish=None):
|
||||||
if batch.forward_mode.is_extend() or batch.is_extend_in_batch:
|
if batch.forward_mode.is_extend() or batch.is_extend_in_batch:
|
||||||
# Target prefill
|
# Target prefill
|
||||||
|
|||||||
@@ -109,10 +109,6 @@ class NGRAMWorker(BaseSpecWorker):
|
|||||||
loaded,
|
loaded,
|
||||||
)
|
)
|
||||||
|
|
||||||
@property
|
|
||||||
def target_worker(self) -> TpModelWorker:
|
|
||||||
return self._target_worker
|
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def draft_worker(self) -> Optional[EagleDraftWorkerBase]:
|
def draft_worker(self) -> Optional[EagleDraftWorkerBase]:
|
||||||
# NGRAM has no draft model; drafts come from the CPU-side corpus.
|
# NGRAM has no draft model; drafts come from the CPU-side corpus.
|
||||||
|
|||||||
@@ -1,10 +1,11 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import contextlib
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import time
|
import time
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from typing import TYPE_CHECKING, List, Optional
|
from typing import TYPE_CHECKING, Any, List, Optional, Tuple
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from huggingface_hub import snapshot_download
|
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
|
from sglang.srt.speculative.eagle_utils import eagle_prepare_for_decode
|
||||||
|
|
||||||
eagle_prepare_for_decode(batch)
|
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
|
import logging
|
||||||
from typing import Optional, Tuple
|
from typing import Optional
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.environ import envs
|
|
||||||
from sglang.srt.layers.moe.utils import speculative_moe_backend_context
|
from sglang.srt.layers.moe.utils import speculative_moe_backend_context
|
||||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
from sglang.srt.managers.tp_worker import TpModelWorker
|
||||||
from sglang.srt.server_args import ServerArgs
|
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_utils import default_tree_mask_mode
|
||||||
from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker, EAGLEWorkerV2
|
from sglang.srt.speculative.eagle_worker_v2 import EagleDraftWorker, EAGLEWorkerV2
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
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
|
from sglang.srt.utils import empty_context, get_bool_env_var, is_cuda
|
||||||
|
|
||||||
if 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")
|
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):
|
class StandaloneDraftWorker(EagleDraftWorker):
|
||||||
"""Custom EagleDraftWorker that doesn't share embeddings/lm_head with target model."""
|
"""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
|
draft_tp_context if server_args.enable_dp_attention else empty_context
|
||||||
)
|
)
|
||||||
self.tree_mask_mode = default_tree_mask_mode()
|
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).
|
# draft_forward reads this (set in EagleDraftWorker.__init__, skipped here).
|
||||||
self.index_share_for_mtp_iteration = (
|
self.index_share_for_mtp_iteration = (
|
||||||
getattr(
|
getattr(
|
||||||
@@ -210,7 +197,7 @@ class StandaloneWorkerV2(EAGLEWorkerV2):
|
|||||||
)
|
)
|
||||||
self.extend_lens = torch.empty((), dtype=torch.int64, device=self.device)
|
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
|
# TODO: Adaptive speculative
|
||||||
self.adaptive_controller: Optional[AdaptiveController] = None
|
self.adaptive_controller: Optional[AdaptiveController] = None
|
||||||
|
|||||||
Reference in New Issue
Block a user