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()
|
embedding_cache_size = envs.SGLANG_VLM_CACHE_SIZE_MB.get()
|
||||||
init_mm_embedding_cache(embedding_cache_size * 1024 * 1024)
|
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
|
"""Return (draft_token_to_kv_pool, draft_model_config) for the current
|
||||||
draft worker, or (None, None) when no draft KV pool is available."""
|
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
|
return None, None
|
||||||
|
|
||||||
if self.spec_algorithm.supports_spec_v2() and self.enable_overlap:
|
if spec_algorithm.supports_spec_v2() and enable_overlap:
|
||||||
if self.server_args.enable_multi_layer_eagle:
|
if server_args.enable_multi_layer_eagle:
|
||||||
draft_runner = self.draft_worker.draft_worker.draft_runner_list[0]
|
draft_runner = draft_worker.draft_worker.draft_runner_list[0]
|
||||||
else:
|
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 draft_runner.token_to_kv_pool, draft_runner.model_config
|
||||||
|
|
||||||
return (
|
return (
|
||||||
self.draft_worker.model_runner.token_to_kv_pool,
|
draft_worker.model_runner.token_to_kv_pool,
|
||||||
self.draft_worker.model_config,
|
draft_worker.model_config,
|
||||||
)
|
)
|
||||||
|
|
||||||
def _maybe_register_hicache_draft(self) -> None:
|
def _maybe_register_hicache_draft(self) -> None:
|
||||||
@@ -1021,7 +1028,12 @@ class Scheduler(
|
|||||||
if not self.enable_hierarchical_cache:
|
if not self.enable_hierarchical_cache:
|
||||||
return
|
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:
|
if draft_kv_pool is None:
|
||||||
return
|
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?
|
# 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
|
# Default to the target model_config so the MetadataBuffers branches
|
||||||
# below can always access it; overridden by the draft model_config
|
# below can always access it; overridden by the draft model_config
|
||||||
# when this node runs a spec module.
|
# when this node runs a spec module.
|
||||||
|
|||||||
Reference in New Issue
Block a user