[NPU]pp support mla kv transfer (#23893)

This commit is contained in:
chenxu214
2026-05-13 09:10:02 +08:00
committed by GitHub
parent d6d3d0f599
commit 642ac9c916
6 changed files with 93 additions and 23 deletions
+65 -18
View File
@@ -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 = [
@@ -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:
@@ -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),
)
@@ -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,
)
+12 -4
View File
@@ -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
@@ -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)