Files
enginex-ascend-910-vllm/vllm_ascend/core/kv_cache_interface.py
Sun Ruoxi 7f8a1b1f7a init v0.23.0
Signed-off-by: Sun Ruoxi <sunruoxi@4paradigm.com>
2026-08-27 15:11:51 +08:00

316 lines
13 KiB
Python

# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
from dataclasses import dataclass, field
import torch
from typing_extensions import Self
from vllm.config import VllmConfig
from vllm.utils.math_utils import cdiv
from vllm.utils.torch_utils import get_dtype_size
from vllm.v1.core.single_type_kv_cache_manager import SlidingWindowManager
from vllm.v1.kv_cache_interface import FullAttentionSpec, MLAAttentionSpec, SlidingWindowMLASpec
from vllm.v1.kv_cache_spec_registry import KVCacheSpecRegistry
from vllm_ascend.core.single_type_kv_cache_manager import CompressAttentionManager
from vllm_ascend.utils import AscendDeviceType, get_ascend_device_type
def _get_c8_k_cache_dtype() -> torch.dtype:
return torch.float8_e4m3fn if get_ascend_device_type() == AscendDeviceType.A5 else torch.int8
def _get_c8_k_scale_cache_dtype() -> torch.dtype:
return torch.float32 if get_ascend_device_type() == AscendDeviceType.A5 else torch.float16
@dataclass(frozen=True, kw_only=True)
class AscendMLAAttentionSpec(MLAAttentionSpec):
"""MLAAttentionSpec extended to support DSA models, with independent SFA and LI C8 support.
When LI C8 is enabled, the KV cache tuple changes from
(kv_cache[0]: bfloat16, kv_cache[1]: bfloat16, kv_cache[2]: bfloat16)
to
(kv_cache[0]: bfloat16, kv_cache[1]: bfloat16, kv_cache[2]: int8, kv_cache[3]: float16).
The semantic meaning of each native KV cache entry is as follows:
1. kv_cache[0] stores kv_lora.
2. kv_cache[1] stores k_rope.
3. kv_cache[2] stores the key tensor from the indexer module.
4. kv_cache[3] stores the key scale tensor from the indexer module,
and exists only when LI C8 is enabled.
With SFA C8, kv_lora, k_rope, and per-tile quantization scales are
packed into kv_cache[0]. The resulting cache is (packed_kv, indexer_k)
or (packed_kv, indexer_k, indexer_scale) when LI C8 is also enabled.
The main changes are as follows:
1. The key tensor from the indexer module stored in kv_cache[2] is
converted from bf16 to int8 to reduce memory usage. It is then
processed with int8 precision in Lightning_indexer computation
to improve computational efficiency.
2. The quantization scale of the key tensor in the indexer module
must also be stored for the Lightning_indexer_quant operator,
and is therefore saved in kv_cache[3].
"""
scale_dim: int = 0
scale_dtype: torch.dtype = torch.int8
sparse_head_dim: tuple[int, ...] | None = None
cache_sparse_sfa_c8: bool = False
cache_sparse_li_c8: bool = False
c8_k_cache_dtype: torch.dtype = field(default_factory=_get_c8_k_cache_dtype)
c8_k_scale_cache_dtype: torch.dtype = field(default_factory=_get_c8_k_scale_cache_dtype)
sfa_dcp_replicated_indexer_size: int = 1
@property
def page_size_bytes(self) -> int:
if self.cache_sparse_sfa_c8:
assert self.sparse_head_dim is not None
assert len(self.sparse_head_dim) == 3
num_heads_per_page = self.block_size * self.num_kv_heads
ckv_head_dim, qk_rope_head_dim, index_head_dim = self.sparse_head_dim
assert qk_rope_head_dim == 0
ckv_bytes = num_heads_per_page * ckv_head_dim * get_dtype_size(self.c8_k_cache_dtype)
qli_dtype = self.c8_k_cache_dtype if self.cache_sparse_li_c8 else self.dtype
qli_bytes = (
num_heads_per_page * index_head_dim * self.sfa_dcp_replicated_indexer_size * get_dtype_size(qli_dtype)
)
qli_scale_bytes = (
num_heads_per_page * self.sfa_dcp_replicated_indexer_size * get_dtype_size(self.c8_k_scale_cache_dtype)
if self.cache_sparse_li_c8 and index_head_dim > 0
else 0
)
return ckv_bytes + qli_bytes + qli_scale_bytes
if self.cache_sparse_li_c8:
assert self.sparse_head_dim is not None
assert len(self.sparse_head_dim) == 3
k_head_dim, v_head_dim, index_head_dim = self.sparse_head_dim
assert index_head_dim > 0
num_heads_per_page = self.block_size * self.num_kv_heads
return num_heads_per_page * (
(k_head_dim + v_head_dim) * get_dtype_size(self.dtype)
+ index_head_dim * self.sfa_dcp_replicated_indexer_size * get_dtype_size(self.c8_k_cache_dtype)
+ self.sfa_dcp_replicated_indexer_size * get_dtype_size(self.c8_k_scale_cache_dtype)
)
if (
self.sparse_head_dim is not None
and len(self.sparse_head_dim) == 3
and self.sfa_dcp_replicated_indexer_size > 1
):
k_head_dim, v_head_dim, index_head_dim = self.sparse_head_dim
replicated_head_size = k_head_dim + v_head_dim + index_head_dim * self.sfa_dcp_replicated_indexer_size
return (
self.block_size
* self.num_kv_heads
* (
replicated_head_size * get_dtype_size(self.dtype)
+ self.scale_dim * get_dtype_size(self.scale_dtype)
)
)
return (
self.block_size
* self.num_kv_heads
* (self.head_size * get_dtype_size(self.dtype) + self.scale_dim * get_dtype_size(self.scale_dtype))
)
@property
def sparse_kv_cache_ratio(self) -> tuple[float, float | None, float | None, float | None]:
"""
Compute the relative byte share of each KV cache entry.
Returns:
A tuple containing the ratios for:
- kv_cache[0]
- kv_cache[1]
- kv_cache[2]
- kv_cache[3] (None if Sparse C8 is disabled or Sparse C8 on A5 device)
"""
assert self.sparse_head_dim is not None
if self.cache_sparse_sfa_c8:
ckv_head_dim, qk_rope_head_dim, index_k_head_dim = self.sparse_head_dim
assert qk_rope_head_dim == 0
ckv_virtual = ckv_head_dim * get_dtype_size(self.c8_k_cache_dtype)
if index_k_head_dim == 0:
return (
1.0,
None,
None,
None,
)
qli_dtype = self.c8_k_cache_dtype if self.cache_sparse_li_c8 else self.dtype
qli_virtual = index_k_head_dim * self.sfa_dcp_replicated_indexer_size * get_dtype_size(qli_dtype)
scale_virtual = (
self.sfa_dcp_replicated_indexer_size * get_dtype_size(self.c8_k_scale_cache_dtype)
if self.cache_sparse_li_c8
else 0
)
total_virtual_head_dim = ckv_virtual + qli_virtual + scale_virtual
return (
total_virtual_head_dim / ckv_virtual,
total_virtual_head_dim / qli_virtual,
total_virtual_head_dim / scale_virtual if scale_virtual > 0 else None,
None,
)
k_head_dim, v_head_dim, index_head_dim = self.sparse_head_dim
replicated_index_head_dim = index_head_dim * self.sfa_dcp_replicated_indexer_size
if self.cache_sparse_li_c8:
k_virtual = k_head_dim * get_dtype_size(self.dtype)
v_virtual = v_head_dim * get_dtype_size(self.dtype)
qli_virtual = replicated_index_head_dim * get_dtype_size(self.c8_k_cache_dtype)
scale_virtual = self.sfa_dcp_replicated_indexer_size * get_dtype_size(self.c8_k_scale_cache_dtype)
total_virtual_head_dim = k_virtual + v_virtual + qli_virtual + scale_virtual
return (
total_virtual_head_dim / k_virtual,
total_virtual_head_dim / v_virtual,
total_virtual_head_dim / qli_virtual,
total_virtual_head_dim / scale_virtual,
)
total_virtual_head_dim = k_head_dim + v_head_dim + replicated_index_head_dim
return (
total_virtual_head_dim / k_head_dim,
total_virtual_head_dim / v_head_dim,
total_virtual_head_dim / replicated_index_head_dim if replicated_index_head_dim > 0 else None,
None,
)
@classmethod
def merge(cls, specs: list[Self]) -> Self:
assert all(isinstance(spec, MLAAttentionSpec) for spec in specs), (
"All attention layers in the same KV cache group must be MLAAttentionSpec."
)
layout_set = {
(
spec.block_size,
spec.num_kv_heads,
spec.head_size,
spec.scale_dim,
spec.scale_dtype,
spec.sparse_head_dim,
spec.dtype,
)
for spec in specs
}
assert len(layout_set) == 1, (
"All attention layers in the same KV cache group must use the same KV cache layout."
)
cache_dtype_str_set = set(spec.cache_dtype_str for spec in specs)
assert len(cache_dtype_str_set) == 1, (
"All attention layers in the same KV cache group must use the same quantization method."
)
cache_sparse_sfa_c8_set = set(spec.cache_sparse_sfa_c8 for spec in specs)
assert len(cache_sparse_sfa_c8_set) == 1, (
"All attention layers in the same KV cache group must use the same sparse SFA C8 setting."
)
cache_sparse_li_c8_set = set(spec.cache_sparse_li_c8 for spec in specs)
assert len(cache_sparse_li_c8_set) == 1, (
"All attention layers in the same KV cache group must use the same sparse LI C8 setting."
)
sfa_dcp_replicated_indexer_size_set = set(spec.sfa_dcp_replicated_indexer_size for spec in specs)
assert len(sfa_dcp_replicated_indexer_size_set) == 1, (
"All attention layers in the same KV cache group must use the same SFA DCP replicated indexer size."
)
return cls(
block_size=specs[0].block_size,
num_kv_heads=specs[0].num_kv_heads,
head_size=specs[0].head_size,
scale_dim=specs[0].scale_dim,
scale_dtype=specs[0].scale_dtype,
sparse_head_dim=specs[0].sparse_head_dim,
dtype=specs[0].dtype,
cache_dtype_str=cache_dtype_str_set.pop(),
cache_sparse_sfa_c8=specs[0].cache_sparse_sfa_c8,
cache_sparse_li_c8=specs[0].cache_sparse_li_c8,
sfa_dcp_replicated_indexer_size=sfa_dcp_replicated_indexer_size_set.pop(),
)
def max_memory_usage_bytes(self, vllm_config: VllmConfig) -> int:
max_model_len = vllm_config.model_config.max_model_len
dcp_world_size = vllm_config.parallel_config.decode_context_parallel_size
pcp_world_size = vllm_config.parallel_config.prefill_context_parallel_size
# Note(hc): each dcp rank only need save
# (max_model_len//dcp_world_size) tokens locally.
if dcp_world_size * pcp_world_size > 1:
max_model_len = cdiv(max_model_len, dcp_world_size * pcp_world_size)
return cdiv(max_model_len, self.block_size * self.compress_ratio) * self.page_size_bytes
@dataclass(frozen=True, kw_only=True)
class AscendSlidingWindowMLASpec(SlidingWindowMLASpec):
"""Sliding window attention with MLA cache format."""
cache_dtype_str: str | None = None
# DeepseekV4-only: see MLAAttentionSpec.model_version.
alignment: int | None = None # Default to None for no padding.
compress_ratio: int = 1
model_version: str | None = None
def __post_init__(self):
pass
@property
def storage_block_size(self) -> int:
return self.block_size
@property
def real_page_size_bytes(self) -> int:
return self.storage_block_size * self.num_kv_heads * self.head_size * get_dtype_size(self.dtype)
@classmethod
def merge(cls, specs: list[Self]) -> Self:
assert all(isinstance(spec, AscendSlidingWindowMLASpec) for spec in specs), (
"All attention layers in the same KV cache group must be AscendSlidingWindowMLASpec."
)
cache_dtype_str_set = set(spec.cache_dtype_str for spec in specs)
compress_ratio_set = set(spec.compress_ratio for spec in specs)
model_version_set = set(spec.model_version for spec in specs)
sliding_window_set = set(spec.sliding_window for spec in specs)
assert (
len(cache_dtype_str_set) == 1
and len(compress_ratio_set) == 1
and len(model_version_set) == 1
and len(sliding_window_set) == 1
), (
"All attention layers in the same KV cache group must use the same "
"quantization method, compress ratio, model version and sliding "
"window size."
)
return cls(
block_size=specs[0].block_size,
num_kv_heads=specs[0].num_kv_heads,
head_size=specs[0].head_size,
dtype=specs[0].dtype,
page_size_padded=specs[0].page_size_padded,
sliding_window=sliding_window_set.pop(),
cache_dtype_str=cache_dtype_str_set.pop(),
compress_ratio=compress_ratio_set.pop(),
model_version=model_version_set.pop(),
)
def register_ascend_kv_cache_specs() -> None:
KVCacheSpecRegistry.register(
kvcache_spec_cls=AscendMLAAttentionSpec,
manager_class=CompressAttentionManager,
uniform_type_base_spec=FullAttentionSpec,
)
KVCacheSpecRegistry.register(
kvcache_spec_cls=AscendSlidingWindowMLASpec,
manager_class=SlidingWindowManager,
uniform_type_base_spec=SlidingWindowMLASpec,
)