Skip to content
135 changes: 135 additions & 0 deletions tests/v1/attention/test_kpool_indexer_block_sizes.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,135 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""CPU tests for storage-block kernel-block selection in prepare_kernel_block_sizes."""

from types import SimpleNamespace

import pytest
import torch

from vllm.config import VllmConfig, set_current_vllm_config
from vllm.models.glm5next.common.attention import Glm5NextIndexerCache
from vllm.v1.attention.backend import AttentionBackend, MultipleOf
from vllm.v1.attention.backends.mla import indexer as indexer_mod
from vllm.v1.attention.backends.mla.indexer import (
DeepseekV32IndexerBackend,
KpoolTailBackend,
)
from vllm.v1.attention.backends.mla.rocm_aiter_mla_sparse import (
ROCMAiterMLASparseBackend,
)
from vllm.v1.kv_cache_interface import MLAAttentionSpec
from vllm.v1.worker.utils import prepare_kernel_block_sizes

pytestmark = pytest.mark.cpu_test

# ``index_kpool`` of zai-org/GLM-5.3-Flash: one indexer state pools 4
# compressed states (hence storage_block_size = page * 4).
INDEX_KPOOL = 4

# The hybrid KDA/mamba manager floors at TP8..TP1; 640 is the geometry of the
# original report ("the 640-token block table").
TP_FLOOR_MANAGER_BLOCKS = [640, 1152, 2176, 4352]

Comment thread
vllmellm marked this conversation as resolved.

def _mock_rocm_platform(monkeypatch: pytest.MonkeyPatch) -> None:
# Platform predicates are mutually exclusive in production. Override both
# so these ROCm policy tests do not inherit CUDA capabilities from the CI
# host running them (same approach as test_deepseek_v4_rocm_adaptive).
monkeypatch.setattr(indexer_mod.current_platform, "is_cuda", lambda: False)
monkeypatch.setattr(indexer_mod.current_platform, "is_rocm", lambda: True)


def _kpool_storage_spec(manager_block: int) -> MLAAttentionSpec:
"""Real Glm5NextIndexerCache spec for the given cache block size."""
with set_current_vllm_config(VllmConfig()):
cache = Glm5NextIndexerCache(
head_dim=128,
dtype=torch.bfloat16,
prefix="model.layers.0.indexer.k_cache_probe",
cache_config=SimpleNamespace(block_size=manager_block),
index_kpool=INDEX_KPOOL,
)
spec = cache.get_kv_cache_spec(VllmConfig())
assert isinstance(spec, MLAAttentionSpec)
return spec


def _prepare(
spec: MLAAttentionSpec, backends: list[type[AttentionBackend]]
) -> list[int]:
kv_cache_config = SimpleNamespace(
kv_cache_groups=[SimpleNamespace(kv_cache_spec=spec)]
)
attn_groups = [[SimpleNamespace(backend=backend) for backend in backends]]
return prepare_kernel_block_sizes(kv_cache_config, attn_groups)


@pytest.mark.parametrize("manager_block", TP_FLOOR_MANAGER_BLOCKS)
def test_prepare_uses_storage_block_for_the_kpool_group(
monkeypatch: pytest.MonkeyPatch, manager_block: int
):
"""The bug: on ROCm, kpool groups must get 128/256 - not the 640-4352
manager block - as their kernel block size.

Manager-granular selection is what made the block table address manager
blocks while the kpool writer and the index-cache gather read pool pages,
aliasing every cache column past the request's row.
"""
_mock_rocm_platform(monkeypatch)
spec = _kpool_storage_spec(manager_block)
expected_storage = 256 if manager_block % 256 == 0 else 128
assert spec.storage_block_size == expected_storage

selected = _prepare(
spec,
[ROCMAiterMLASparseBackend, DeepseekV32IndexerBackend, KpoolTailBackend],
)
assert selected == [expected_storage]
assert selected != [manager_block]
# The hybrid block-table split stays integral.
assert manager_block % selected[0] == 0


def test_prepare_falls_back_when_storage_block_is_not_supported():
"""No storage block, or a storage block the backends reject (e.g. the
CUDA indexer's exact ``[64]`` vs a 128-token storage block): selection is
`select_common_block_size`'s answer, unchanged from no-fix behavior.
"""

class Fixed64Backend:
@staticmethod
def get_supported_kernel_block_sizes():
return [64]

@staticmethod
def get_name() -> str:
return "FIXED64"

# Storage block the group's backend does not accept -> backend vote (64).
spec = _kpool_storage_spec(640)
assert spec.storage_block_size == 128
assert _prepare(spec, [Fixed64Backend]) == [64]

# No storage block at all (any non-kpool MLA model) -> backend vote.
plain_spec = MLAAttentionSpec(
block_size=640,
num_kv_heads=1,
head_size=128,
dtype=torch.bfloat16,
)
assert plain_spec.storage_block_size is None
assert _prepare(plain_spec, [Fixed64Backend]) == [64]

# ...and with backends that accept the manager block, the manager wins,
# exactly as on main.
class AcceptAllBackend:
@staticmethod
def get_supported_kernel_block_sizes():
return [MultipleOf(1)]

@staticmethod
def get_name() -> str:
return "ACCEPTALL"

assert _prepare(plain_spec, [AcceptAllBackend]) == [640]
67 changes: 43 additions & 24 deletions vllm/v1/worker/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -327,6 +327,30 @@ def update_draft_decode_metadata(
self.get_metadata_builder().update_draft_decode_metadata(metadata)


def _block_size_is_supported(
backends: list[type[AttentionBackend]], block_size: int
) -> bool:
"""Check if the block size is supported by all backends.

An exact ``int`` declaration must match exactly; a ``MultipleOf``
declaration accepts any multiple of its base.
"""
for backend in backends:
is_supported = False
for supported_size in backend.get_supported_kernel_block_sizes():
if isinstance(supported_size, int):
if block_size == supported_size:
is_supported = True
elif isinstance(supported_size, MultipleOf):
if block_size % supported_size.base == 0:
is_supported = True
else:
raise ValueError(f"Unknown supported size: {supported_size}")
if not is_supported:
return False
return True


def select_common_block_size(
kv_manager_block_size: int,
backends: list[type[AttentionBackend]],
Expand All @@ -348,27 +372,7 @@ def select_common_block_size(
ValueError: If no valid block size found.

"""

def block_size_is_supported(
backends: list[type[AttentionBackend]], block_size: int
) -> bool:
"""Check if the block size is supported by all backends."""
for backend in backends:
is_supported = False
for supported_size in backend.get_supported_kernel_block_sizes():
if isinstance(supported_size, int):
if block_size == supported_size:
is_supported = True
elif isinstance(supported_size, MultipleOf):
if block_size % supported_size.base == 0:
is_supported = True
else:
raise ValueError(f"Unknown supported size: {supported_size}")
if not is_supported:
return False
return True

if block_size_is_supported(backends, kv_manager_block_size):
if _block_size_is_supported(backends, kv_manager_block_size):
return kv_manager_block_size

# MultipleOf constraints also accept the manager size if they accept a divisor.
Expand All @@ -381,7 +385,7 @@ def block_size_is_supported(
}

for size in sorted(candidates, reverse=True):
if block_size_is_supported(backends, size):
if _block_size_is_supported(backends, size):
return size
raise ValueError(
f"No common block size for {kv_manager_block_size} ("
Expand Down Expand Up @@ -493,9 +497,24 @@ def prepare_kernel_block_sizes(
# This is an attention backend that supports virtual block splitting.
kv_manager_block_size = kv_cache_group.kv_cache_spec.block_size
group_backends = [g.backend for g in attn_groups[kv_cache_gid]]
selected_kernel_size = select_common_block_size(
kv_manager_block_size, group_backends
storage_block_size = (
kv_cache_spec.storage_block_size
if isinstance(kv_cache_spec, MLAAttentionSpec)
else None
)
if storage_block_size is not None and _block_size_is_supported(
group_backends, storage_block_size
):
# Storage-block specs (e.g. the GLM-5.3-Flash kpool indexer
# cache) address the cache in pool pages, and every other
# consumer (cache views, metadata builders, hisparse) already
# uses storage_block_size as the kernel block. Fall back to the
# backend vote when the group's backends do not accept it.
selected_kernel_size = storage_block_size
else:
selected_kernel_size = select_common_block_size(
kv_manager_block_size, group_backends
)
kernel_block_sizes.append(selected_kernel_size)
elif isinstance(kv_cache_spec, MambaSpec):
# This is likely Mamba or other non-attention cache, no splitting.
Expand Down
Loading