Decouple _get_draft_kv_pool from self before extraction (#25601)
This commit is contained in:
@@ -998,22 +998,29 @@ class Scheduler(
|
||||
embedding_cache_size = envs.SGLANG_VLM_CACHE_SIZE_MB.get()
|
||||
init_mm_embedding_cache(embedding_cache_size * 1024 * 1024)
|
||||
|
||||
def _get_draft_kv_pool(self):
|
||||
@staticmethod
|
||||
def get_draft_kv_pool(
|
||||
*,
|
||||
draft_worker: "BaseTpWorker",
|
||||
spec_algorithm: SpeculativeAlgorithm,
|
||||
server_args: ServerArgs,
|
||||
enable_overlap: bool,
|
||||
):
|
||||
"""Return (draft_token_to_kv_pool, draft_model_config) for the current
|
||||
draft worker, or (None, None) when no draft KV pool is available."""
|
||||
if self.draft_worker is None or self.spec_algorithm.is_ngram():
|
||||
if draft_worker is None or spec_algorithm.is_ngram():
|
||||
return None, None
|
||||
|
||||
if self.spec_algorithm.supports_spec_v2() and self.enable_overlap:
|
||||
if self.server_args.enable_multi_layer_eagle:
|
||||
draft_runner = self.draft_worker.draft_worker.draft_runner_list[0]
|
||||
if spec_algorithm.supports_spec_v2() and enable_overlap:
|
||||
if server_args.enable_multi_layer_eagle:
|
||||
draft_runner = draft_worker.draft_worker.draft_runner_list[0]
|
||||
else:
|
||||
draft_runner = self.draft_worker.draft_worker.draft_runner
|
||||
draft_runner = draft_worker.draft_worker.draft_runner
|
||||
return draft_runner.token_to_kv_pool, draft_runner.model_config
|
||||
|
||||
return (
|
||||
self.draft_worker.model_runner.token_to_kv_pool,
|
||||
self.draft_worker.model_config,
|
||||
draft_worker.model_runner.token_to_kv_pool,
|
||||
draft_worker.model_config,
|
||||
)
|
||||
|
||||
def _maybe_register_hicache_draft(self) -> None:
|
||||
@@ -1021,7 +1028,12 @@ class Scheduler(
|
||||
if not self.enable_hierarchical_cache:
|
||||
return
|
||||
|
||||
draft_kv_pool, _ = self._get_draft_kv_pool()
|
||||
draft_kv_pool, _ = Scheduler.get_draft_kv_pool(
|
||||
draft_worker=self.draft_worker,
|
||||
spec_algorithm=self.spec_algorithm,
|
||||
server_args=self.server_args,
|
||||
enable_overlap=self.enable_overlap,
|
||||
)
|
||||
if draft_kv_pool is None:
|
||||
return
|
||||
|
||||
@@ -1216,7 +1228,12 @@ class Scheduler(
|
||||
)
|
||||
|
||||
# todo: should we fix this when enabling mtp or it doesn't matter since we only enable mtp in decode node thus we don't transfer draft kvs between P and D?
|
||||
draft_token_to_kv_pool, model_config = self._get_draft_kv_pool()
|
||||
draft_token_to_kv_pool, model_config = Scheduler.get_draft_kv_pool(
|
||||
draft_worker=self.draft_worker,
|
||||
spec_algorithm=self.spec_algorithm,
|
||||
server_args=self.server_args,
|
||||
enable_overlap=self.enable_overlap,
|
||||
)
|
||||
# Default to the target model_config so the MetadataBuffers branches
|
||||
# below can always access it; overridden by the draft model_config
|
||||
# when this node runs a spec module.
|
||||
|
||||
Reference in New Issue
Block a user