Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions tests/ut/attention/test_sfa_nope_metadata.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@ def _builder(block_size, a5, monkeypatch, rope_dim=0):
get_topk_lengths=lambda positions: torch.where(positions == 0, 1, 7),
)
config = SimpleNamespace(
cache_config=SimpleNamespace(block_size=block_size),
model_config=SimpleNamespace(
max_model_len=4096,
get_head_size=lambda: 512,
Expand Down
15 changes: 10 additions & 5 deletions tests/ut/models/test_glm5next_kv_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -171,6 +171,7 @@ def test_indexer_metadata_addresses_complete_storage_pages(storage_block_size):
logical_size = storage_block_size * pool_size
split = logical_size // 128
config = SimpleNamespace(
cache_config=SimpleNamespace(block_size=logical_size),
scheduler_config=SimpleNamespace(max_num_batched_tokens=4, max_num_seqs=1),
model_config=SimpleNamespace(max_model_len=logical_size * 3),
)
Expand Down Expand Up @@ -203,17 +204,20 @@ def test_indexer_metadata_addresses_complete_storage_pages(storage_block_size):
block_table_tensor=expanded,
)
first, draft = [builder.build(0, common) for builder in builders]
assert first.block_size == storage_block_size
torch.testing.assert_close(first.block_table, pages)
# Kernel-granularity blocks: the metadata reports the natural kernel rows
# (128 tokens / pool ratio) and passes the common expanded table through
# as a view, so both builders observe the same persistent buffer.
assert first.block_size == 128 // pool_size
torch.testing.assert_close(first.block_table, expanded)
assert first.slot_mapping.tolist() == [8 * storage_block_size - 1, -1, 2 * storage_block_size, -1]
assert first.seq_lens.tolist() == [storage_block_size + 1]
address = first.block_table.data_ptr()
assert draft.block_table.data_ptr() != address
assert draft.block_table.data_ptr() == address
common.block_table_tensor[:, :split] = 3 * split + torch.arange(split)
refreshed = builders[0].build(0, common)
assert refreshed.block_table.data_ptr() == address
assert first.block_table.tolist() == [[3, 2, -1]]
assert draft.block_table.tolist() == [[7, 2, -1]]
torch.testing.assert_close(refreshed.block_table, common.block_table_tensor[:1])
torch.testing.assert_close(draft.block_table, common.block_table_tensor[:1])


def test_model_cache_layers_publish_source_compatible_specs():
Expand Down Expand Up @@ -265,6 +269,7 @@ def test_model_cache_layers_publish_source_compatible_specs():

def test_indexer_metadata_preserves_raw_request_boundaries():
config = SimpleNamespace(
cache_config=SimpleNamespace(block_size=256),
scheduler_config=SimpleNamespace(
max_num_batched_tokens=16,
max_num_seqs=2,
Expand Down
3 changes: 3 additions & 0 deletions tests/ut/worker/test_attn_utils_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -999,6 +999,9 @@ def _make_mla_layer(*, fa_quant: bool = False, sparse_c8: bool = False):
head_size=128,
dtype=torch.bfloat16,
cache_dtype_str="auto",
model_version=None,
non_causal_multi_token_decode=False,
**({"compress_ratio": 1} if vllm_version_is("0.28.0") else {"tokens_per_state": 1}),
)
return layer

Expand Down
49 changes: 8 additions & 41 deletions vllm_ascend/attention/indexer_kpool.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,6 @@
from vllm.config import VllmConfig
from vllm.config.compilation import CUDAGraphMode
from vllm.forward_context import get_forward_context
from vllm.utils.math_utils import cdiv
from vllm.v1.attention.backend import (
AttentionBackend,
AttentionCGSupport,
Expand Down Expand Up @@ -82,7 +81,7 @@ def __init__(
if not layer_names or any(not name.endswith(".indexer.k_cache") for name in layer_names):
raise ValueError(f"Invalid Indexer KPool cache layer names: {layer_names}.")
super().__init__(kv_cache_spec, layer_names, vllm_config, device)
self.logical_block_size = kv_cache_spec.block_size
self.logical_block_size = vllm_config.cache_config.block_size
self.storage_block_size = get_storage_block_size(kv_cache_spec)
if self.storage_block_size <= 0:
raise ValueError(f"Indexer KPool storage block size must be positive, got {self.storage_block_size}.")
Expand All @@ -93,7 +92,7 @@ def __init__(
f"kernel block size: logical={self.logical_block_size}, "
f"kernel={GLM5_NEXT_SFA_KERNEL_BLOCK_SIZE}."
)
self.kernel_blocks_per_logical_block = self.logical_block_size // GLM5_NEXT_SFA_KERNEL_BLOCK_SIZE
self.kernel_row_block_size = GLM5_NEXT_SFA_KERNEL_BLOCK_SIZE // self.compress_ratio
scheduler_config = vllm_config.scheduler_config
# ACLGraph replay keeps the addresses captured on the first run. The
# derived compressed metadata therefore needs persistent storage that
Expand All @@ -118,16 +117,6 @@ def __init__(
dtype=torch.int32,
device=device,
)
max_logical_blocks = cdiv(
vllm_config.model_config.max_model_len,
self.logical_block_size,
)
self._block_table_buffer = torch.empty(
scheduler_config.max_num_seqs,
max_logical_blocks,
dtype=torch.int32,
device=device,
)

def build(
self,
Expand Down Expand Up @@ -168,38 +157,14 @@ def build(
seq_lens_cpu = None
if seq_lens_cpu is not None:
seq_lens_cpu = torch.div(seq_lens_cpu, self.compress_ratio, rounding_mode="floor")
expanded_block_table = common_attn_metadata.block_table_tensor[:num_reqs]
split = self.kernel_blocks_per_logical_block
if expanded_block_table.shape[1] % split:
raise ValueError(
"GLM-Next indexer received a partially expanded SFA block "
f"table: width={expanded_block_table.shape[1]}, split={split}."
)
logical_width = expanded_block_table.shape[1] // split
if logical_width > self._block_table_buffer.shape[1]:
raise ValueError(
"GLM-Next indexer block table exceeds its persistent buffer: "
f"required={logical_width}, capacity="
f"{self._block_table_buffer.shape[1]}."
)
block_table = self._block_table_buffer[:num_reqs, :logical_width]
# The common full-group table is expanded for the C128 SFA kernel:
# scheduler block N becomes [split*N, ..., split*N+split-1]. The
# compressed indexer owns one physical page per scheduler block, so it
# must recover N rather than treating the SFA sub-blocks as pages.
torch.div(
expanded_block_table[:, ::split],
split,
rounding_mode="floor",
out=block_table,
)
block_table = common_attn_metadata.block_table_tensor[:num_reqs]
return AscendIndexerKPoolMetadata(
block_table=block_table,
slot_mapping=slot_mapping,
seq_lens=seq_lens,
seq_lens_cpu=seq_lens_cpu,
positions=positions,
block_size=self.storage_block_size,
block_size=self.kernel_row_block_size,
compress_ratio=self.compress_ratio,
cum_query_lens=cum_query_lens,
raw_seq_lens=raw_seq_lens,
Expand Down Expand Up @@ -235,8 +200,9 @@ def get_kv_cache_shape(
num_kv_heads: int,
head_size: int,
cache_type: str = "",
cache_dtype_str: str = "auto",
) -> tuple[int, ...]:
del cache_type
del cache_type, cache_dtype_str
if num_kv_heads != 1:
raise ValueError(f"Indexer KPool cache requires one KV head, got {num_kv_heads}.")
return (num_blocks, block_size, num_kv_heads, head_size)
Expand Down Expand Up @@ -324,8 +290,9 @@ def get_kv_cache_shape(
num_kv_heads: int,
head_size: int,
cache_type: str = "",
cache_dtype_str: str = "auto",
) -> tuple[int, ...]:
del cache_type
del cache_type, cache_dtype_str
if num_kv_heads != 1:
raise ValueError(f"Indexer KPool tail cache requires one KV head, got {num_kv_heads}.")
return (num_blocks, 2, block_size, head_size)
Expand Down
4 changes: 3 additions & 1 deletion vllm_ascend/attention/sfa_v1.py
Original file line number Diff line number Diff line change
Expand Up @@ -180,7 +180,9 @@ def __init__(self, kv_cache_spec, vllm_config, device, indexer, kernel_block_siz
if not self.use_smla and block_size > SPARSE_ATTENTION_MAX_BLOCK_SIZE:
self.block_size = kernel_block_size
self.table_stride = self.block_size // kernel_block_size
table_width = cdiv(vllm_config.model_config.max_model_len, block_size) * (block_size // self.block_size)
cache_block_size = vllm_config.cache_config.block_size
expand_factor = max(cache_block_size // kernel_block_size, 1)
table_width = cdiv(vllm_config.model_config.max_model_len, cache_block_size) * expand_factor
self.block_table_buffer = torch.empty(
vllm_config.scheduler_config.max_num_seqs,
table_width,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@
SlidingWindowSpec,
)

from vllm_ascend.core.kv_cache_interface import AscendSFAIndexerCacheSpec
from vllm_ascend.core.kv_cache_interface import AscendIndexerKPoolTailSpec, AscendSFAIndexerCacheSpec
from vllm_ascend.distributed.kv_transfer.kv_p2p.mooncake.base_worker import (
MooncakeBaseConnectorWorker,
)
Expand Down Expand Up @@ -310,7 +310,7 @@ def _get_layer_remote_tp_rank_groups(
return self._get_mamba_remote_tp_rank_groups(remote_tp_size)

fixed_total_num_kv_heads = None
if isinstance(spec, (AscendSFAIndexerCacheSpec, SlidingWindowMLASpec)):
if isinstance(spec, (AscendSFAIndexerCacheSpec, AscendIndexerKPoolTailSpec, SlidingWindowMLASpec)):
local_dcp_size = remote_dcp_size = 1
local_num_kv_heads = remote_num_kv_heads = 1
fixed_total_num_kv_heads = 1
Expand Down Expand Up @@ -995,6 +995,8 @@ def _append_block_transfer_addresses(
remote_inner_offset: int = 0,
) -> None:
"""Append addresses from block groups prepared for one request."""
if transfer_len <= 0:
return
for local_block_group, remote_block_group in zip(local_block_groups, remote_block_groups):
src_list.append(local_base_addr + local_block_group[0] * local_block_stride + local_inner_offset)
dst_list.append(remote_base_addr + remote_block_group[0] * remote_block_stride + remote_inner_offset)
Expand Down Expand Up @@ -1058,7 +1060,7 @@ def _append_spec_transfer_addresses(
remote_tp_metadata = remote_metadata.metadata_by_tp_rank[remote_tp_rank]
transfer_whole_block = isinstance(
spec,
(MLAAttentionSpec, SlidingWindowMLASpec, AscendSFAIndexerCacheSpec),
(MLAAttentionSpec, SlidingWindowMLASpec, AscendSFAIndexerCacheSpec, AscendIndexerKPoolTailSpec),
)
if transfer_whole_block:
for (local_layer_index, remote_layer_index), transfer_entries in transfer_entries_by_layer.items():
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -111,7 +111,7 @@ def merge_storage_range(

cache_tensors: list[torch.Tensor] = []
for layer_name in layer_names:
cache_tensors.extend(as_kv_cache_tensors(kv_caches.get(layer_name)))
cache_tensors.extend(cache for cache in as_kv_cache_tensors(kv_caches.get(layer_name)) if cache.numel() > 0)

caches_by_storage: dict[int, list[torch.Tensor]] = {}
for cache in cache_tensors:
Expand Down
64 changes: 61 additions & 3 deletions vllm_ascend/worker/v2/attn_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -55,6 +55,7 @@
AscendMLAAttentionSpec,
AscendSFAIndexerCacheSpec,
AscendSlidingWindowMLASpec,
get_kv_cache_compression_ratio,
get_storage_block_size,
)
from vllm_ascend.device.hardware_profile import HardwareCapability, get_current_hardware_profile
Expand Down Expand Up @@ -154,13 +155,28 @@ def get_kv_cache_spec(vllm_config: VllmConfig) -> dict[str, KVCacheSpec]:
head_size = spec.head_size
dtype = spec.dtype
cache_dtype_str = spec.cache_dtype_str
model_version = spec.model_version or getattr(attn_module, "model_version", None)
indexes_kv_by_block_stride = bool(
getattr(spec, "indexes_kv_by_block_stride", False)
or getattr(attn_module, "indexes_kv_by_block_stride", False)
)
compression_ratio = get_kv_cache_compression_ratio(spec)
ratio_kwargs: dict[str, Any] = (
{"compress_ratio": compression_ratio}
if vllm_version_is("0.28.0")
else {"tokens_per_state": compression_ratio}
)
spec = AscendMLAAttentionSpec(
block_size=spec.block_size,
num_kv_heads=spec.num_kv_heads,
head_size=head_size,
dtype=dtype,
cache_dtype_str=cache_dtype_str,
cache_sparse_sfa_c8=cache_sparse_sfa_c8,
non_causal_multi_token_decode=spec.non_causal_multi_token_decode,
Comment thread
sunbaosong marked this conversation as resolved.
model_version=model_version,
indexes_kv_by_block_stride=indexes_kv_by_block_stride,
**ratio_kwargs,
)
if isinstance(attn_module, DeepseekV32IndexerCache):
if not getattr(
Expand Down Expand Up @@ -1033,6 +1049,11 @@ def _reshape_kv_cache_v2(
if group_storage_block_size != group_spec.block_size
else kernel_block_sizes[group.kv_cache_group_id]
)
if group_storage_block_size != group_spec.block_size and getattr(
group_spec, "indexes_kv_by_block_stride", False
):
compression_ratio = get_kv_cache_compression_ratio(group_spec)
kernel_block_size = kernel_block_sizes[group.kv_cache_group_id] // compression_ratio

for layer_name in group.layer_names:
if layer_name in shared_kv_cache_layers:
Expand Down Expand Up @@ -1127,9 +1148,27 @@ def _reshape_kv_cache_v2(
kv_caches[layer_name] = typed_cache.view(kv_cache_shape)
continue

if isinstance(kv_cache_spec, AscendIndexerKPoolTailSpec) or (
is_dsv4_model and isinstance(kv_cache_spec, (AscendMLAAttentionSpec, AscendSlidingWindowMLASpec))
):
if isinstance(kv_cache_spec, AscendIndexerKPoolTailSpec):
if not isinstance(raw_cache, torch.Tensor):
raise ValueError(f"KPool tail cache for {layer_name} must use one raw tensor.")
typed_slot = raw_cache.view(kv_cache_spec.dtype)
tail_block_el = kv_cache_spec.unpadded_page_size_bytes // get_dtype_size(kv_cache_spec.dtype)
num_tail_blocks = kv_cache_config.num_blocks
if num_tail_blocks * tail_block_el * 2 > typed_slot.numel():
raise ValueError(
f"KPool tail cache for {layer_name} exceeds half the small slot: "
f"packed={num_tail_blocks * tail_block_el} elements, slot={typed_slot.numel()}."
)
kv_caches[layer_name] = [
typed_slot[typed_slot.numel() - num_tail_blocks * tail_block_el :].view(
num_tail_blocks,
2,
kv_cache_spec.block_size,
kv_cache_spec.head_size,
)
]
continue
if is_dsv4_model and isinstance(kv_cache_spec, (AscendMLAAttentionSpec, AscendSlidingWindowMLASpec)):
if not isinstance(raw_cache, torch.Tensor):
raise ValueError(f"DSA cache for {layer_name} must use one raw tensor.")
kv_caches[layer_name] = _view_dsv4_cache(
Expand Down Expand Up @@ -1172,6 +1211,25 @@ def _reshape_kv_cache_v2(
cache_dtype,
)
sparse_sfa_c8 = enable_sfa(vllm_config) and bool(getattr(kv_cache_spec, "cache_sparse_sfa_c8", False))
if isinstance(kv_cache_spec, (AscendMLAAttentionSpec, MLAAttentionSpec)) and (
get_kv_cache_compression_ratio(kv_cache_spec) > 1
):
raw_single = raw_cache[0] if isinstance(raw_cache, tuple) else raw_cache
if isinstance(raw_cache, tuple) and len(raw_cache) != 1:
raise ValueError(f"Compressed indexer cache for {layer_name} must be a single tensor.")
shape = tuple(kv_cache_shape)
strides = [1] * len(shape)
for dim_idx in range(len(shape) - 2, -1, -1):
strides[dim_idx] = strides[dim_idx + 1] * shape[dim_idx + 1]
typed_slot = raw_single.view(kv_cache_spec.dtype)
if strides[0] * shape[0] * 2 > typed_slot.numel():
raise ValueError(
f"Compressed indexer cache for {layer_name} exceeds half the small slot: "
f"packed={strides[0] * shape[0]} elements, slot={typed_slot.numel()}."
)
cache = torch.as_strided(typed_slot, size=shape, stride=tuple(strides))
kv_caches[layer_name] = (cache,)
continue
if isinstance(kv_cache_spec, (AscendMLAAttentionSpec, MLAAttentionSpec)):
num_blocks_, block_size_, num_kv_heads, _ = kv_cache_shape
k_dim, v_dim = _get_attention_kv_cache_dims(layer_name, kv_cache_spec)
Expand Down
Loading