From 642ac9c916d533de253f69ffe43799f5c42b4da5 Mon Sep 17 00:00:00 2001 From: chenxu214 Date: Wed, 13 May 2026 09:10:02 +0800 Subject: [PATCH] [NPU]pp support mla kv transfer (#23893) --- .../sglang/srt/disaggregation/ascend/conn.py | 83 +++++++++++++++---- python/sglang/srt/disaggregation/base/conn.py | 4 + python/sglang/srt/disaggregation/decode.py | 1 + python/sglang/srt/disaggregation/prefill.py | 1 + python/sglang/srt/disaggregation/utils.py | 16 +++- .../hardware_backend/npu/memory_pool_npu.py | 11 ++- 6 files changed, 93 insertions(+), 23 deletions(-) diff --git a/python/sglang/srt/disaggregation/ascend/conn.py b/python/sglang/srt/disaggregation/ascend/conn.py index 098b905da..90b020846 100644 --- a/python/sglang/srt/disaggregation/ascend/conn.py +++ b/python/sglang/srt/disaggregation/ascend/conn.py @@ -41,6 +41,36 @@ class AscendKVManager(MooncakeKVManager): ): self.engine.batch_register(component_ptrs, component_lens) + def get_mla_kv_ptrs_with_pp( + self, src_kv_ptrs: List[int], dst_kv_ptrs: List[int] + ) -> Tuple[List[int], List[int], int]: + # src_kv_ptrs: k_data, v_data, index_k_data(optional) + # dst_kv_ptrs: k_data, v_data, index_k_data(optional) + start_layer = self.kv_args.prefill_start_layer + kv_buf_groups = getattr(self.kv_args, "kv_buf_groups", 1) + total_kv_layers = getattr(self.kv_args, "total_kv_layers", 0) + src_layers = len(src_kv_ptrs) // kv_buf_groups + # When only speculative-algorithm is enabled for decode + # the KV has one more layer than prefill. + # The draft layer needs to be skipped. + dst_total_layers = ( + min(len(dst_kv_ptrs) // kv_buf_groups, total_kv_layers) + if total_kv_layers + else len(dst_kv_ptrs) // kv_buf_groups + ) + end_layer = start_layer + src_layers + if src_layers == dst_total_layers: + sliced_dst_kv_ptrs = dst_kv_ptrs + else: + sliced_dst_kv_ptrs = [] + for i in range(kv_buf_groups): + layer_offset = i * dst_total_layers + sliced_dst_kv_ptrs.extend( + dst_kv_ptrs[layer_offset + start_layer : layer_offset + end_layer] + ) + layers_current_pp_stage = len(src_kv_ptrs) + return src_kv_ptrs, sliced_dst_kv_ptrs, layers_current_pp_stage + def send_kvcache( self, mooncake_session_id: str, @@ -55,25 +85,42 @@ class AscendKVManager(MooncakeKVManager): ) if self.pp_size > 1: - src_k_ptrs, src_v_ptrs, dst_k_ptrs, dst_v_ptrs, layers_current_pp_stage = ( - self.get_mha_kv_ptrs_with_pp(self.kv_args.kv_data_ptrs, dst_kv_ptrs) - ) + if self.is_mla_backend: + src_kv_ptrs, sliced_dst_kv_ptrs, layers_current_pp_stage = ( + self.get_mla_kv_ptrs_with_pp(self.kv_args.kv_data_ptrs, dst_kv_ptrs) + ) + layers_params = [ + ( + src_kv_ptrs[layer_id], + sliced_dst_kv_ptrs[layer_id], + self.kv_args.kv_item_lens[layer_id], + ) + for layer_id in range(layers_current_pp_stage) + ] + else: + ( + src_k_ptrs, + src_v_ptrs, + dst_k_ptrs, + dst_v_ptrs, + layers_current_pp_stage, + ) = self.get_mha_kv_ptrs_with_pp(self.kv_args.kv_data_ptrs, dst_kv_ptrs) - layers_params = [ - ( - src_k_ptrs[layer_id], - dst_k_ptrs[layer_id], - self.kv_args.kv_item_lens[layer_id], - ) - for layer_id in range(layers_current_pp_stage) - ] + [ - ( - src_v_ptrs[layer_id], - dst_v_ptrs[layer_id], - self.kv_args.kv_item_lens[layers_current_pp_stage + layer_id], - ) - for layer_id in range(layers_current_pp_stage) - ] + layers_params = [ + ( + src_k_ptrs[layer_id], + dst_k_ptrs[layer_id], + self.kv_args.kv_item_lens[layer_id], + ) + for layer_id in range(layers_current_pp_stage) + ] + [ + ( + src_v_ptrs[layer_id], + dst_v_ptrs[layer_id], + self.kv_args.kv_item_lens[layers_current_pp_stage + layer_id], + ) + for layer_id in range(layers_current_pp_stage) + ] else: num_layers = len(self.kv_args.kv_data_ptrs) layers_params = [ diff --git a/python/sglang/srt/disaggregation/base/conn.py b/python/sglang/srt/disaggregation/base/conn.py index edc44895d..264f3ca0d 100644 --- a/python/sglang/srt/disaggregation/base/conn.py +++ b/python/sglang/srt/disaggregation/base/conn.py @@ -52,6 +52,10 @@ class KVArgs: prefill_start_layer: int # for system dp system_dp_rank: int + # Only used of npu, for kv buf groups + kv_buf_groups: int + # Only used of npu, for decode total kv layers + total_kv_layers: int class KVPoll: diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index b717256cf..1b0a654ce 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -411,6 +411,7 @@ class DecodePreallocQueue: kv_args, self.token_to_kv_pool, self.draft_token_to_kv_pool, + total_kv_layers=self.scheduler.model_config.num_hidden_layers, req_to_token_pool=getattr(self, "req_to_token_pool", None), ) diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index 62f033ec3..7ddcbe169 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -181,6 +181,7 @@ class PrefillBootstrapQueue: kv_args, self.token_to_kv_pool, self.draft_token_to_kv_pool, + self.scheduler.model_config.num_hidden_layers, req_to_token_pool=req_to_token_pool, ) diff --git a/python/sglang/srt/disaggregation/utils.py b/python/sglang/srt/disaggregation/utils.py index 2a0b4beec..951fa5b7d 100644 --- a/python/sglang/srt/disaggregation/utils.py +++ b/python/sglang/srt/disaggregation/utils.py @@ -553,6 +553,7 @@ def setup_state_kv_args( kv_args: KVArgs, token_to_kv_pool, draft_token_to_kv_pool=None, + total_kv_layers: int = None, req_to_token_pool=None, ) -> None: """Populate ``kv_args`` state-buffer fields from the given pool. @@ -560,6 +561,7 @@ def setup_state_kv_args( lives in one place. """ from sglang.srt.disaggregation.base.conn import StateType + from sglang.srt.hardware_backend.npu.memory_pool_npu import NPUMLATokenToKVPool from sglang.srt.mem_cache.base_swa_memory_pool import BaseSWAKVPool from sglang.srt.mem_cache.memory_pool import HybridLinearKVPool, NSATokenToKVPool @@ -587,7 +589,7 @@ def setup_state_kv_args( append_state_component( kv_args, StateType.MAMBA, data_ptrs, data_lens, item_lens, dim ) - elif isinstance(token_to_kv_pool, NSATokenToKVPool): + elif isinstance(token_to_kv_pool, (NSATokenToKVPool, NPUMLATokenToKVPool)): if draft_token_to_kv_pool is not None and isinstance( draft_token_to_kv_pool, NSATokenToKVPool ): @@ -599,9 +601,15 @@ def setup_state_kv_args( data_ptrs = data_ptrs + draft_data_ptrs data_lens = data_lens + draft_data_lens item_lens = item_lens + draft_item_lens - append_state_component( - kv_args, StateType.NSA, data_ptrs, data_lens, item_lens - ) + if isinstance(token_to_kv_pool, NPUMLATokenToKVPool): + kv_args.kv_buf_groups = ( + len(kv_args.kv_data_ptrs) // token_to_kv_pool.layer_num + ) + kv_args.total_kv_layers = total_kv_layers + else: + append_state_component( + kv_args, StateType.NSA, data_ptrs, data_lens, item_lens + ) if ( StateType.MAMBA not in kv_args.state_types diff --git a/python/sglang/srt/hardware_backend/npu/memory_pool_npu.py b/python/sglang/srt/hardware_backend/npu/memory_pool_npu.py index e4f319fa5..a1a9a0cdb 100644 --- a/python/sglang/srt/hardware_backend/npu/memory_pool_npu.py +++ b/python/sglang/srt/hardware_backend/npu/memory_pool_npu.py @@ -1,7 +1,6 @@ from typing import TYPE_CHECKING, Optional import torch -import torch_npu from sglang.srt.constants import GPU_MEMORY_TYPE_KV_CACHE from sglang.srt.mem_cache.memory_pool import ( @@ -10,10 +9,14 @@ from sglang.srt.mem_cache.memory_pool import ( get_tensor_size_bytes, ) from sglang.srt.utils import get_bool_env_var +from sglang.srt.utils.common import is_npu if TYPE_CHECKING: from sglang.srt.layers.radix_attention import RadixAttention +if is_npu(): + import torch_npu + def _init_npu_conv_state( conv_state_in, conv_state_shape, speculative_num_draft_tokens: Optional[int] = None @@ -292,6 +295,12 @@ class NPUMLATokenToKVPool(MLATokenToKVPool): self.v_buffer[layer_id - self.start_layer], ) + def get_state_buf_infos(self): + data_ptrs = [self.index_k_buffer[i].data_ptr() for i in range(self.layer_num)] + data_lens = [self.index_k_buffer[i].nbytes for i in range(self.layer_num)] + item_lens = [self.index_k_buffer[i][0].nbytes for i in range(self.layer_num)] + return data_ptrs, data_lens, item_lens + def get_key_buffer(self, layer_id: int): if self.layer_transfer_counter is not None: self.layer_transfer_counter.wait_until(layer_id - self.start_layer)