dedup state_kv_args setup into helper (#24340)
This commit is contained in:
@@ -47,6 +47,7 @@ from sglang.srt.disaggregation.utils import (
|
||||
poll_and_all_reduce,
|
||||
poll_and_all_reduce_with_staging,
|
||||
prepare_abort,
|
||||
setup_state_kv_args,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
||||
@@ -367,44 +368,7 @@ class DecodePreallocQueue:
|
||||
self.metadata_buffers.get_buf_infos()
|
||||
)
|
||||
|
||||
if hasattr(self.token_to_kv_pool, "get_state_buf_infos"):
|
||||
state_data_ptrs, state_data_lens, state_item_lens = (
|
||||
self.token_to_kv_pool.get_state_buf_infos()
|
||||
)
|
||||
kv_args.state_data_ptrs = state_data_ptrs
|
||||
kv_args.state_data_lens = state_data_lens
|
||||
kv_args.state_item_lens = state_item_lens
|
||||
|
||||
if isinstance(self.token_to_kv_pool, SWAKVPool):
|
||||
kv_args.state_type = "swa"
|
||||
elif isinstance(self.token_to_kv_pool, HybridLinearKVPool):
|
||||
kv_args.state_type = "mamba"
|
||||
# Get state dimension info for cross-TP slice transfer
|
||||
if hasattr(self.token_to_kv_pool, "get_state_dim_per_tensor"):
|
||||
kv_args.state_dim_per_tensor = (
|
||||
self.token_to_kv_pool.get_state_dim_per_tensor()
|
||||
)
|
||||
elif isinstance(self.token_to_kv_pool, NSATokenToKVPool):
|
||||
kv_args.state_type = "nsa"
|
||||
if self.draft_token_to_kv_pool is not None and isinstance(
|
||||
self.draft_token_to_kv_pool, NSATokenToKVPool
|
||||
):
|
||||
(
|
||||
draft_state_data_ptrs,
|
||||
draft_state_data_lens,
|
||||
draft_state_item_lens,
|
||||
) = self.draft_token_to_kv_pool.get_state_buf_infos()
|
||||
kv_args.state_data_ptrs += draft_state_data_ptrs
|
||||
kv_args.state_data_lens += draft_state_data_lens
|
||||
kv_args.state_item_lens += draft_state_item_lens
|
||||
|
||||
else:
|
||||
kv_args.state_type = "none"
|
||||
else:
|
||||
kv_args.state_data_ptrs = []
|
||||
kv_args.state_data_lens = []
|
||||
kv_args.state_item_lens = []
|
||||
kv_args.state_type = "none"
|
||||
setup_state_kv_args(kv_args, self.token_to_kv_pool, self.draft_token_to_kv_pool)
|
||||
|
||||
kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device
|
||||
kv_args.gpu_id = self.scheduler.gpu_id
|
||||
|
||||
@@ -39,6 +39,7 @@ from sglang.srt.disaggregation.utils import (
|
||||
is_mla_backend,
|
||||
poll_and_all_reduce_attn_cp_tp_group,
|
||||
prepare_abort,
|
||||
setup_state_kv_args,
|
||||
)
|
||||
from sglang.srt.environ import envs
|
||||
from sglang.srt.managers.schedule_batch import (
|
||||
@@ -176,44 +177,7 @@ class PrefillBootstrapQueue:
|
||||
kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device
|
||||
kv_args.gpu_id = self.scheduler.gpu_id
|
||||
|
||||
if hasattr(self.token_to_kv_pool, "get_state_buf_infos"):
|
||||
state_data_ptrs, state_data_lens, state_item_lens = (
|
||||
self.token_to_kv_pool.get_state_buf_infos()
|
||||
)
|
||||
kv_args.state_data_ptrs = state_data_ptrs
|
||||
kv_args.state_data_lens = state_data_lens
|
||||
kv_args.state_item_lens = state_item_lens
|
||||
|
||||
if isinstance(self.token_to_kv_pool, SWAKVPool):
|
||||
kv_args.state_type = "swa"
|
||||
elif isinstance(self.token_to_kv_pool, HybridLinearKVPool):
|
||||
kv_args.state_type = "mamba"
|
||||
# Get state dimension info for cross-TP slice transfer
|
||||
if hasattr(self.token_to_kv_pool, "get_state_dim_per_tensor"):
|
||||
kv_args.state_dim_per_tensor = (
|
||||
self.token_to_kv_pool.get_state_dim_per_tensor()
|
||||
)
|
||||
elif isinstance(self.token_to_kv_pool, NSATokenToKVPool):
|
||||
kv_args.state_type = "nsa"
|
||||
if self.draft_token_to_kv_pool is not None and isinstance(
|
||||
self.draft_token_to_kv_pool, NSATokenToKVPool
|
||||
):
|
||||
(
|
||||
draft_state_data_ptrs,
|
||||
draft_state_data_lens,
|
||||
draft_state_item_lens,
|
||||
) = self.draft_token_to_kv_pool.get_state_buf_infos()
|
||||
kv_args.state_data_ptrs += draft_state_data_ptrs
|
||||
kv_args.state_data_lens += draft_state_data_lens
|
||||
kv_args.state_item_lens += draft_state_item_lens
|
||||
|
||||
else:
|
||||
kv_args.state_type = "none"
|
||||
else:
|
||||
kv_args.state_data_ptrs = []
|
||||
kv_args.state_data_lens = []
|
||||
kv_args.state_item_lens = []
|
||||
kv_args.state_type = "none"
|
||||
setup_state_kv_args(kv_args, self.token_to_kv_pool, self.draft_token_to_kv_pool)
|
||||
|
||||
kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
|
||||
kv_manager = kv_manager_class(
|
||||
|
||||
@@ -531,6 +531,57 @@ def is_mla_backend(target_kv_pool) -> bool:
|
||||
return isinstance(target_kv_pool, MLATokenToKVPool)
|
||||
|
||||
|
||||
def setup_state_kv_args(
|
||||
kv_args: KVArgs,
|
||||
token_to_kv_pool,
|
||||
draft_token_to_kv_pool=None,
|
||||
) -> None:
|
||||
"""Populate ``kv_args`` state-buffer fields from the given pool.
|
||||
|
||||
Shared by prefill and decode bootstrap paths so the state_type dispatch
|
||||
lives in one place.
|
||||
"""
|
||||
from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool, NSATokenToKVPool
|
||||
from sglang.srt.mem_cache.swa_memory_pool import SWAKVPool
|
||||
|
||||
if not hasattr(token_to_kv_pool, "get_state_buf_infos"):
|
||||
kv_args.state_data_ptrs = []
|
||||
kv_args.state_data_lens = []
|
||||
kv_args.state_item_lens = []
|
||||
kv_args.state_type = "none"
|
||||
return
|
||||
|
||||
state_data_ptrs, state_data_lens, state_item_lens = (
|
||||
token_to_kv_pool.get_state_buf_infos()
|
||||
)
|
||||
kv_args.state_data_ptrs = state_data_ptrs
|
||||
kv_args.state_data_lens = state_data_lens
|
||||
kv_args.state_item_lens = state_item_lens
|
||||
|
||||
if isinstance(token_to_kv_pool, SWAKVPool):
|
||||
kv_args.state_type = "swa"
|
||||
elif isinstance(token_to_kv_pool, HybridLinearKVPool):
|
||||
kv_args.state_type = "mamba"
|
||||
# Get state dimension info for cross-TP slice transfer
|
||||
if hasattr(token_to_kv_pool, "get_state_dim_per_tensor"):
|
||||
kv_args.state_dim_per_tensor = token_to_kv_pool.get_state_dim_per_tensor()
|
||||
elif isinstance(token_to_kv_pool, NSATokenToKVPool):
|
||||
kv_args.state_type = "nsa"
|
||||
if draft_token_to_kv_pool is not None and isinstance(
|
||||
draft_token_to_kv_pool, NSATokenToKVPool
|
||||
):
|
||||
(
|
||||
draft_state_data_ptrs,
|
||||
draft_state_data_lens,
|
||||
draft_state_item_lens,
|
||||
) = draft_token_to_kv_pool.get_state_buf_infos()
|
||||
kv_args.state_data_ptrs += draft_state_data_ptrs
|
||||
kv_args.state_data_lens += draft_state_data_lens
|
||||
kv_args.state_item_lens += draft_state_item_lens
|
||||
else:
|
||||
kv_args.state_type = "none"
|
||||
|
||||
|
||||
def prepare_abort(req: Req, error_message: str, status_code=None):
|
||||
from sglang.srt.managers.schedule_batch import FINISH_ABORT
|
||||
|
||||
|
||||
Reference in New Issue
Block a user