Skip to content
Closed
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
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
from vllm.v1.kv_cache_interface import (
AttentionSpec,
MambaSpec,
MLAAttentionSpec,
UniformTypeKVCacheSpecs,
)
from vllm.v1.kv_offload.base import (
Expand Down Expand Up @@ -52,6 +53,11 @@ def register_kv_caches(
):
kv_cache_config = self.spec.kv_cache_config
num_blocks = kv_cache_config.num_blocks
model_config = self.spec.vllm_config.model_config
parallel_config = self.spec.vllm_config.parallel_config
total_num_kv_heads = model_config.get_total_num_kv_heads()
tp_size = parallel_config.tensor_parallel_size
dcp_size = parallel_config.decode_context_parallel_size

# Packed layouts (e.g. DSv4) set block_stride > 0; their tensors use
# stride(0) as the manager-block stride (equals total_num_bytes_per_block).
Expand All @@ -69,6 +75,8 @@ def register_kv_caches(
unpadded_page_size_bytes: dict[str, int] = {}
# layer_name -> size of page in bytes
page_size_bytes: dict[str, int] = {}
# layer_name -> canonical (n_heads, h_stride, bs_stride, replicated)
head_layout: dict[str, tuple[int, int, int, bool]] = {}
for kv_cache_group in kv_cache_config.kv_cache_groups:
group_layer_names = kv_cache_group.layer_names
group_kv_cache_spec = kv_cache_group.kv_cache_spec
Expand Down Expand Up @@ -108,6 +116,34 @@ def register_kv_caches(
unpadded_page_size_bytes[layer_name] = (
layer_kv_cache_spec.real_page_size_bytes
)
if isinstance(layer_kv_cache_spec, MLAAttentionSpec):
# Replicated latent: any one rank's page is complete
# (unless DCP shards tokens across ranks).
if dcp_size == 1:
head_layout[layer_name] = (0, 0, 0, True)
else:
page = unpadded_page_size_bytes[layer_name]
num_head_cells = (
layer_kv_cache_spec.block_size
* layer_kv_cache_spec.num_kv_heads
)
h_stride = page // num_head_cells
# Head-sliceable only if TP ranks hold distinct heads
# in rank order and the page is pure head cells;
# otherwise fail closed to the opaque default.
if (
dcp_size == 1
and layer_kv_cache_spec.num_kv_heads * tp_size
== total_num_kv_heads
and not layer_kv_cache_spec.kv_quant_mode.is_per_token_head
and h_stride * num_head_cells == page
):
head_layout[layer_name] = (
total_num_kv_heads,
h_stride,
total_num_kv_heads * h_stride,
False,
)

elif isinstance(layer_kv_cache_spec, MambaSpec):
state_tensors = kv_caches[layer_name]
Expand Down Expand Up @@ -200,10 +236,17 @@ def register_kv_caches(

curr_tensor_idx = len(block_tensors) - 1
for layer_name in tensor_layer_names:
n_heads, h_stride, bs_stride, replicated = head_layout.get(
layer_name, (0, 0, 0, False)
)
block_data_refs[layer_name].append(
CanonicalKVCacheRef(
tensor_idx=curr_tensor_idx,
page_size_bytes=(unpadded_page_size_bytes[layer_name]),
n_heads=n_heads,
h_stride=h_stride,
bs_stride=bs_stride,
replicated=replicated,
)
)

Expand Down
9 changes: 9 additions & 0 deletions vllm/v1/kv_offload/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -415,6 +415,15 @@ class CanonicalKVCacheRef:
tensor_idx: int
# The un-padded page size per block in bytes
page_size_bytes: int
# Global KV head count (across all TP ranks); 0 = no head decomposition
n_heads: int = 0
# Bytes per head per token (the head-slice unit)
h_stride: int = 0
# Bytes per token (== n_heads * h_stride)
bs_stride: int = 0
# When n_heads == 0: True = page is identical on every TP rank (MLA
# latent), False = rank-specific shards (Mamba states, packed blocks)
replicated: bool = False


@dataclass
Expand Down
Loading