[NPU] Support Hybrid KV Cache for Ascend backend (#18032)
Co-authored-by: gengjinsong <gengjinsong@huawei.com>
This commit is contained in:
@@ -51,6 +51,9 @@ class ForwardMetadata:
|
|||||||
# calculated map for kv positions [bs * maxseqlen]
|
# calculated map for kv positions [bs * maxseqlen]
|
||||||
block_tables: Optional[torch.Tensor] = None
|
block_tables: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
|
# mapped block_tables for swa
|
||||||
|
block_tables_swa: Optional[torch.Tensor] = None
|
||||||
|
|
||||||
# seq len inputs
|
# seq len inputs
|
||||||
extend_seq_lens_cpu_int: Optional[torch.Tensor] = None
|
extend_seq_lens_cpu_int: Optional[torch.Tensor] = None
|
||||||
seq_lens_cpu_int: Optional[torch.Tensor] = None
|
seq_lens_cpu_int: Optional[torch.Tensor] = None
|
||||||
@@ -252,6 +255,11 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
)
|
)
|
||||||
if self.use_mla:
|
if self.use_mla:
|
||||||
self.ringmla_mask = self.ascend_attn_mask_builder.ringmla_mask
|
self.ringmla_mask = self.ascend_attn_mask_builder.ringmla_mask
|
||||||
|
self.is_hybrid_swa = model_runner.is_hybrid_swa
|
||||||
|
if self.is_hybrid_swa:
|
||||||
|
self.full_to_swa_index_mapping = (
|
||||||
|
model_runner.token_to_kv_pool.full_to_swa_index_mapping
|
||||||
|
)
|
||||||
|
|
||||||
# dllm model config
|
# dllm model config
|
||||||
self.dllm_config = DllmConfig.from_server_args(model_runner.server_args)
|
self.dllm_config = DllmConfig.from_server_args(model_runner.server_args)
|
||||||
@@ -285,6 +293,19 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
][:, :: self.page_size]
|
][:, :: self.page_size]
|
||||||
// self.page_size
|
// self.page_size
|
||||||
)
|
)
|
||||||
|
if self.is_hybrid_swa:
|
||||||
|
self.forward_metadata.block_tables_swa = (
|
||||||
|
(
|
||||||
|
self.full_to_swa_index_mapping[
|
||||||
|
forward_batch.req_to_token_pool.req_to_token[
|
||||||
|
forward_batch.req_pool_indices, :seq_lens_max
|
||||||
|
]
|
||||||
|
][:, :: self.page_size]
|
||||||
|
// self.page_size
|
||||||
|
)
|
||||||
|
.to(torch.int32)
|
||||||
|
.contiguous()
|
||||||
|
)
|
||||||
if forward_batch.extend_seq_lens is not None:
|
if forward_batch.extend_seq_lens is not None:
|
||||||
self.forward_metadata.extend_seq_lens = forward_batch.extend_seq_lens
|
self.forward_metadata.extend_seq_lens = forward_batch.extend_seq_lens
|
||||||
self.forward_metadata.extend_seq_lens_cpu_int = (
|
self.forward_metadata.extend_seq_lens_cpu_int = (
|
||||||
@@ -347,6 +368,12 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
device=self.device,
|
device=self.device,
|
||||||
),
|
),
|
||||||
}
|
}
|
||||||
|
if self.is_hybrid_swa:
|
||||||
|
self.graph_metadata["block_tables_swa"] = torch.empty(
|
||||||
|
(max_bs, (self.max_context_len + self.page_size - 1) // self.page_size),
|
||||||
|
dtype=torch.int32,
|
||||||
|
device=self.device,
|
||||||
|
)
|
||||||
|
|
||||||
def init_forward_metadata_capture_cuda_graph(
|
def init_forward_metadata_capture_cuda_graph(
|
||||||
self,
|
self,
|
||||||
@@ -375,6 +402,8 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
metadata.block_tables[:bs, max_seq_pages:].fill_(0)
|
metadata.block_tables[:bs, max_seq_pages:].fill_(0)
|
||||||
metadata.block_tables[bs:, :].fill_(0)
|
metadata.block_tables[bs:, :].fill_(0)
|
||||||
|
|
||||||
|
if self.is_hybrid_swa:
|
||||||
|
metadata.block_tables_swa = self.graph_metadata["block_tables_swa"][:bs, :]
|
||||||
metadata.seq_lens_cpu_list = seq_lens.cpu().int().tolist()
|
metadata.seq_lens_cpu_list = seq_lens.cpu().int().tolist()
|
||||||
metadata.seq_lens = seq_lens
|
metadata.seq_lens = seq_lens
|
||||||
if (
|
if (
|
||||||
@@ -428,6 +457,15 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
max_len += self.speculative_num_draft_tokens
|
max_len += self.speculative_num_draft_tokens
|
||||||
max_seq_pages = (max_len + self.page_size - 1) // self.page_size
|
max_seq_pages = (max_len + self.page_size - 1) // self.page_size
|
||||||
|
|
||||||
|
if self.is_hybrid_swa:
|
||||||
|
metadata.block_tables_swa[:bs, :max_seq_pages].copy_(
|
||||||
|
self.full_to_swa_index_mapping[
|
||||||
|
self.req_to_token[req_pool_indices[:bs], :max_len]
|
||||||
|
][:, :: self.page_size]
|
||||||
|
// self.page_size
|
||||||
|
)
|
||||||
|
metadata.block_tables_swa[:bs, max_seq_pages:].fill_(0)
|
||||||
|
metadata.block_tables_swa[bs:, :].fill_(0)
|
||||||
metadata.block_tables[:bs, :max_seq_pages].copy_(
|
metadata.block_tables[:bs, :max_seq_pages].copy_(
|
||||||
self.req_to_token[req_pool_indices[:bs], :max_len][:, :: self.page_size]
|
self.req_to_token[req_pool_indices[:bs], :max_len][:, :: self.page_size]
|
||||||
// self.page_size
|
// self.page_size
|
||||||
@@ -816,13 +854,18 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||||
|
|
||||||
if sinks is not None:
|
if sinks is not None:
|
||||||
|
# Use SWA block tables if hybrid SWA is enabled for this layer
|
||||||
|
if self.is_hybrid_swa and layer.sliding_window_size != -1:
|
||||||
|
block_tables = self.forward_metadata.block_tables_swa
|
||||||
|
else:
|
||||||
|
block_tables = self.forward_metadata.block_tables
|
||||||
attn_out = attention_sinks_prefill_triton(
|
attn_out = attention_sinks_prefill_triton(
|
||||||
q,
|
q,
|
||||||
k_cache,
|
k_cache,
|
||||||
v_cache,
|
v_cache,
|
||||||
sinks,
|
sinks,
|
||||||
self.forward_metadata.extend_seq_lens,
|
self.forward_metadata.extend_seq_lens,
|
||||||
self.forward_metadata.block_tables,
|
block_tables,
|
||||||
self.forward_metadata.seq_lens,
|
self.forward_metadata.seq_lens,
|
||||||
layer.scaling,
|
layer.scaling,
|
||||||
layer.sliding_window_size,
|
layer.sliding_window_size,
|
||||||
@@ -1395,12 +1438,17 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
k_cache = forward_batch.token_to_kv_pool.get_key_buffer(layer.layer_id)
|
||||||
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||||
|
|
||||||
|
# Use SWA block tables if hybrid SWA is enabled for this layer
|
||||||
|
if self.is_hybrid_swa and layer.sliding_window_size != -1:
|
||||||
|
block_tables = self.forward_metadata.block_tables_swa
|
||||||
|
else:
|
||||||
|
block_tables = self.forward_metadata.block_tables
|
||||||
attn_out = attention_sinks_triton(
|
attn_out = attention_sinks_triton(
|
||||||
q,
|
q,
|
||||||
k_cache,
|
k_cache,
|
||||||
v_cache,
|
v_cache,
|
||||||
sinks,
|
sinks,
|
||||||
self.forward_metadata.block_tables,
|
block_tables,
|
||||||
self.forward_metadata.seq_lens,
|
self.forward_metadata.seq_lens,
|
||||||
layer.scaling,
|
layer.scaling,
|
||||||
layer.sliding_window_size,
|
layer.sliding_window_size,
|
||||||
@@ -1584,12 +1632,17 @@ class AscendAttnBackend(AttentionBackend):
|
|||||||
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
v_cache = forward_batch.token_to_kv_pool.get_value_buffer(layer.layer_id)
|
||||||
|
|
||||||
if sinks is not None:
|
if sinks is not None:
|
||||||
|
# Use SWA block tables if hybrid SWA is enabled for this layer
|
||||||
|
if self.is_hybrid_swa and layer.sliding_window_size != -1:
|
||||||
|
block_tables = self.forward_metadata.block_tables_swa
|
||||||
|
else:
|
||||||
|
block_tables = self.forward_metadata.block_tables
|
||||||
attn_out = attention_sinks_triton(
|
attn_out = attention_sinks_triton(
|
||||||
q,
|
q,
|
||||||
k_cache,
|
k_cache,
|
||||||
v_cache,
|
v_cache,
|
||||||
sinks,
|
sinks,
|
||||||
self.forward_metadata.block_tables,
|
block_tables,
|
||||||
self.forward_metadata.seq_lens,
|
self.forward_metadata.seq_lens,
|
||||||
layer.scaling,
|
layer.scaling,
|
||||||
layer.sliding_window_size,
|
layer.sliding_window_size,
|
||||||
|
|||||||
@@ -12,6 +12,14 @@ from sglang.srt.mem_cache.allocator import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.mem_cache.memory_pool import KVCache, MHATokenToKVPool
|
from sglang.srt.mem_cache.memory_pool import KVCache, MHATokenToKVPool
|
||||||
from sglang.srt.mem_cache.utils import maybe_init_custom_mem_pool
|
from sglang.srt.mem_cache.utils import maybe_init_custom_mem_pool
|
||||||
|
from sglang.srt.utils import is_npu
|
||||||
|
|
||||||
|
_is_npu = is_npu()
|
||||||
|
|
||||||
|
if _is_npu:
|
||||||
|
from sglang.srt.hardware_backend.npu.allocator_npu import (
|
||||||
|
NPUPagedTokenToKVPoolAllocator,
|
||||||
|
)
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
GB = 1024 * 1024 * 1024
|
GB = 1024 * 1024 * 1024
|
||||||
@@ -256,7 +264,11 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
need_sort,
|
need_sort,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.full_attn_allocator = PagedTokenToKVPoolAllocator(
|
if _is_npu:
|
||||||
|
PagedTokenToKVPoolAllocatorClass = NPUPagedTokenToKVPoolAllocator
|
||||||
|
else:
|
||||||
|
PagedTokenToKVPoolAllocatorClass = PagedTokenToKVPoolAllocator
|
||||||
|
self.full_attn_allocator = PagedTokenToKVPoolAllocatorClass(
|
||||||
size,
|
size,
|
||||||
page_size,
|
page_size,
|
||||||
dtype,
|
dtype,
|
||||||
@@ -264,7 +276,7 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
kvcache.full_kv_pool,
|
kvcache.full_kv_pool,
|
||||||
need_sort,
|
need_sort,
|
||||||
)
|
)
|
||||||
self.swa_attn_allocator = PagedTokenToKVPoolAllocator(
|
self.swa_attn_allocator = PagedTokenToKVPoolAllocatorClass(
|
||||||
size_swa,
|
size_swa,
|
||||||
page_size,
|
page_size,
|
||||||
dtype,
|
dtype,
|
||||||
@@ -347,7 +359,12 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
assert alloc_full_indices is not None
|
assert alloc_full_indices is not None
|
||||||
assert alloc_swa_indices is not None
|
assert alloc_swa_indices is not None
|
||||||
|
|
||||||
self.full_to_swa_index_mapping[alloc_full_indices] = alloc_swa_indices
|
if _is_npu:
|
||||||
|
self.full_to_swa_index_mapping[alloc_full_indices.to(torch.int64)] = (
|
||||||
|
alloc_swa_indices.to(torch.int64)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.full_to_swa_index_mapping[alloc_full_indices] = alloc_swa_indices
|
||||||
return alloc_full_indices
|
return alloc_full_indices
|
||||||
|
|
||||||
def alloc_extend(
|
def alloc_extend(
|
||||||
@@ -387,7 +404,12 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
assert alloc_full_indices is not None
|
assert alloc_full_indices is not None
|
||||||
assert alloc_swa_indices is not None
|
assert alloc_swa_indices is not None
|
||||||
|
|
||||||
self.full_to_swa_index_mapping[alloc_full_indices] = alloc_swa_indices
|
if _is_npu:
|
||||||
|
self.full_to_swa_index_mapping[alloc_full_indices.to(torch.int64)] = (
|
||||||
|
alloc_swa_indices.to(torch.int64)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.full_to_swa_index_mapping[alloc_full_indices] = alloc_swa_indices
|
||||||
|
|
||||||
return alloc_full_indices
|
return alloc_full_indices
|
||||||
|
|
||||||
@@ -410,7 +432,12 @@ class SWATokenToKVPoolAllocator(BaseTokenToKVPoolAllocator):
|
|||||||
if alloc_full_indices is None or alloc_swa_indices is None:
|
if alloc_full_indices is None or alloc_swa_indices is None:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
self.full_to_swa_index_mapping[alloc_full_indices] = alloc_swa_indices
|
if _is_npu:
|
||||||
|
self.full_to_swa_index_mapping[alloc_full_indices.to(torch.int64)] = (
|
||||||
|
alloc_swa_indices.to(torch.int64)
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
self.full_to_swa_index_mapping[alloc_full_indices] = alloc_swa_indices
|
||||||
|
|
||||||
return alloc_full_indices
|
return alloc_full_indices
|
||||||
|
|
||||||
|
|||||||
@@ -445,7 +445,40 @@ class ModelRunnerKVCacheMixin:
|
|||||||
# Initialize token_to_kv_pool
|
# Initialize token_to_kv_pool
|
||||||
is_nsa_model = is_deepseek_nsa(self.model_config.hf_config)
|
is_nsa_model = is_deepseek_nsa(self.model_config.hf_config)
|
||||||
if self.server_args.attention_backend == "ascend" and not self.mambaish_config:
|
if self.server_args.attention_backend == "ascend" and not self.mambaish_config:
|
||||||
if self.use_mla_backend:
|
if self.is_hybrid_swa:
|
||||||
|
from sglang.srt.hardware_backend.npu.memory_pool_npu import (
|
||||||
|
NPUMHATokenToKVPool,
|
||||||
|
)
|
||||||
|
|
||||||
|
kwargs = {}
|
||||||
|
if self.is_hybrid_swa_compress:
|
||||||
|
kwargs = {
|
||||||
|
"swa_head_num": max(
|
||||||
|
1,
|
||||||
|
self.model_config.hf_text_config.swa_num_key_value_heads
|
||||||
|
// get_attention_tp_size(),
|
||||||
|
),
|
||||||
|
"swa_head_dim": self.model_config.hf_text_config.swa_head_dim,
|
||||||
|
"swa_v_head_dim": self.model_config.hf_text_config.swa_v_head_dim,
|
||||||
|
"v_head_dim": self.model_config.hf_text_config.v_head_dim,
|
||||||
|
}
|
||||||
|
self.token_to_kv_pool = SWAKVPool(
|
||||||
|
size=self.full_max_total_num_tokens,
|
||||||
|
size_swa=self.swa_max_total_num_tokens,
|
||||||
|
page_size=self.page_size,
|
||||||
|
dtype=self.kv_cache_dtype,
|
||||||
|
head_num=self.model_config.get_num_kv_heads(
|
||||||
|
get_attention_tp_size()
|
||||||
|
),
|
||||||
|
head_dim=self.model_config.head_dim,
|
||||||
|
swa_attention_layer_ids=self.model_config.swa_attention_layer_ids,
|
||||||
|
full_attention_layer_ids=self.model_config.full_attention_layer_ids,
|
||||||
|
enable_kvcache_transpose=False,
|
||||||
|
device=self.device,
|
||||||
|
token_to_kv_pool_class=NPUMHATokenToKVPool,
|
||||||
|
**kwargs,
|
||||||
|
)
|
||||||
|
elif self.use_mla_backend:
|
||||||
from sglang.srt.hardware_backend.npu.memory_pool_npu import (
|
from sglang.srt.hardware_backend.npu.memory_pool_npu import (
|
||||||
NPUMLATokenToKVPool,
|
NPUMLATokenToKVPool,
|
||||||
)
|
)
|
||||||
@@ -659,18 +692,29 @@ class ModelRunnerKVCacheMixin:
|
|||||||
self.server_args.attention_backend == "ascend"
|
self.server_args.attention_backend == "ascend"
|
||||||
or self.hybrid_gdn_config is not None
|
or self.hybrid_gdn_config is not None
|
||||||
):
|
):
|
||||||
from sglang.srt.hardware_backend.npu.allocator_npu import (
|
if self.is_hybrid_swa:
|
||||||
NPUPagedTokenToKVPoolAllocator,
|
self.token_to_kv_pool_allocator = SWATokenToKVPoolAllocator(
|
||||||
)
|
self.full_max_total_num_tokens,
|
||||||
|
self.swa_max_total_num_tokens,
|
||||||
|
page_size=self.page_size,
|
||||||
|
dtype=self.kv_cache_dtype,
|
||||||
|
device=self.device,
|
||||||
|
kvcache=self.token_to_kv_pool,
|
||||||
|
need_sort=need_sort,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
from sglang.srt.hardware_backend.npu.allocator_npu import (
|
||||||
|
NPUPagedTokenToKVPoolAllocator,
|
||||||
|
)
|
||||||
|
|
||||||
self.token_to_kv_pool_allocator = NPUPagedTokenToKVPoolAllocator(
|
self.token_to_kv_pool_allocator = NPUPagedTokenToKVPoolAllocator(
|
||||||
self.max_total_num_tokens,
|
self.max_total_num_tokens,
|
||||||
page_size=self.page_size,
|
page_size=self.page_size,
|
||||||
dtype=self.kv_cache_dtype,
|
dtype=self.kv_cache_dtype,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
kvcache=self.token_to_kv_pool,
|
kvcache=self.token_to_kv_pool,
|
||||||
need_sort=need_sort,
|
need_sort=need_sort,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
if self.is_hybrid_swa:
|
if self.is_hybrid_swa:
|
||||||
self.token_to_kv_pool_allocator = SWATokenToKVPoolAllocator(
|
self.token_to_kv_pool_allocator = SWATokenToKVPoolAllocator(
|
||||||
|
|||||||
Reference in New Issue
Block a user