[NPU]pp support mla kv transfer (#23893)
This commit is contained in:
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user