diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 8644e8192..ae666e35a 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -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 diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 140ed72a9..533265ae7 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -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( diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py index 43c323058..0bd5b5b76 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -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