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,
|
||||||
poll_and_all_reduce_with_staging,
|
poll_and_all_reduce_with_staging,
|
||||||
prepare_abort,
|
prepare_abort,
|
||||||
|
setup_state_kv_args,
|
||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
from sglang.srt.layers.dp_attention import get_attention_tp_size
|
||||||
@@ -367,44 +368,7 @@ class DecodePreallocQueue:
|
|||||||
self.metadata_buffers.get_buf_infos()
|
self.metadata_buffers.get_buf_infos()
|
||||||
)
|
)
|
||||||
|
|
||||||
if hasattr(self.token_to_kv_pool, "get_state_buf_infos"):
|
setup_state_kv_args(kv_args, self.token_to_kv_pool, self.draft_token_to_kv_pool)
|
||||||
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"
|
|
||||||
|
|
||||||
kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device
|
kv_args.ib_device = self.scheduler.server_args.disaggregation_ib_device
|
||||||
kv_args.gpu_id = self.scheduler.gpu_id
|
kv_args.gpu_id = self.scheduler.gpu_id
|
||||||
|
|||||||
@@ -39,6 +39,7 @@ from sglang.srt.disaggregation.utils import (
|
|||||||
is_mla_backend,
|
is_mla_backend,
|
||||||
poll_and_all_reduce_attn_cp_tp_group,
|
poll_and_all_reduce_attn_cp_tp_group,
|
||||||
prepare_abort,
|
prepare_abort,
|
||||||
|
setup_state_kv_args,
|
||||||
)
|
)
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.managers.schedule_batch import (
|
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.ib_device = self.scheduler.server_args.disaggregation_ib_device
|
||||||
kv_args.gpu_id = self.scheduler.gpu_id
|
kv_args.gpu_id = self.scheduler.gpu_id
|
||||||
|
|
||||||
if hasattr(self.token_to_kv_pool, "get_state_buf_infos"):
|
setup_state_kv_args(kv_args, self.token_to_kv_pool, self.draft_token_to_kv_pool)
|
||||||
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"
|
|
||||||
|
|
||||||
kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
|
kv_manager_class = get_kv_class(self.transfer_backend, KVClassType.MANAGER)
|
||||||
kv_manager = kv_manager_class(
|
kv_manager = kv_manager_class(
|
||||||
|
|||||||
@@ -531,6 +531,57 @@ def is_mla_backend(target_kv_pool) -> bool:
|
|||||||
return isinstance(target_kv_pool, MLATokenToKVPool)
|
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):
|
def prepare_abort(req: Req, error_message: str, status_code=None):
|
||||||
from sglang.srt.managers.schedule_batch import FINISH_ABORT
|
from sglang.srt.managers.schedule_batch import FINISH_ABORT
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user