# 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, )