Skip to content
Open
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
36 changes: 36 additions & 0 deletions tests/v1/attention/test_dsv4_kernel_block_size.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,36 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project

from vllm.models.deepseek_v4.nvidia.flashinfer_sparse import (
DeepseekV4FlashInferMLASparseBackend,
)
from vllm.models.deepseek_v4.sparse_mla import (
DeepseekV4SparseMLABackend,
dsv4_supported_kernel_block_sizes,
)
from vllm.platforms import current_platform
from vllm.v1.attention.backends.mla.indexer import DeepseekV4IndexerBackend


def test_dsv4_kernel_block_size_sm12x(monkeypatch):
monkeypatch.setattr(
current_platform, "is_device_capability_family", lambda family: family == 120
)
assert dsv4_supported_kernel_block_sizes() == [64]
assert DeepseekV4SparseMLABackend.get_supported_kernel_block_sizes() == [64]
assert DeepseekV4FlashInferMLASparseBackend.get_supported_kernel_block_sizes() == [
64
]
assert DeepseekV4IndexerBackend.get_supported_kernel_block_sizes() == [64]


def test_dsv4_kernel_block_size_not_sm12x(monkeypatch):
monkeypatch.setattr(
current_platform, "is_device_capability_family", lambda family: False
)
assert dsv4_supported_kernel_block_sizes() == [256]
assert DeepseekV4SparseMLABackend.get_supported_kernel_block_sizes() == [256]
assert DeepseekV4FlashInferMLASparseBackend.get_supported_kernel_block_sizes() == [
256
]
assert DeepseekV4IndexerBackend.get_supported_kernel_block_sizes() == [256]
6 changes: 1 addition & 5 deletions vllm/models/deepseek_v4/nvidia/flashinfer_sparse.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@
)
from vllm.platforms.interface import DeviceCapability
from vllm.utils.flashinfer import flashinfer_trtllm_batch_decode_sparse_mla_dsv4
from vllm.v1.attention.backend import AttentionCGSupport, MultipleOf
from vllm.v1.attention.backend import AttentionCGSupport
from vllm.v1.attention.backends.mla.compressor_utils import (
get_dspark_swa_index_width,
)
Expand Down Expand Up @@ -110,10 +110,6 @@ class DeepseekV4FlashInferMLASparseBackend(DeepseekV4SparseMLABackend):
"fp8_ds_mla",
]

@staticmethod
def get_supported_kernel_block_sizes() -> list[int | MultipleOf]:
return [256]

@staticmethod
def get_name() -> str:
return "FLASHINFER_MLA_SPARSE_DSV4"
Expand Down
15 changes: 14 additions & 1 deletion vllm/models/deepseek_v4/sparse_mla.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,19 @@
_C128A_TOPK_ALIGNMENT = 128


def dsv4_supported_kernel_block_sizes() -> list[int | MultipleOf]:
"""Kernel page size for DeepSeek-V4 sparse MLA.

FlashInfer SM120 DSV4 decode is compiled for 64-token pages
(``_DECODE_DSV4_PAGE_BLOCK_SIZE = 64``). Keep manager ``--block-size 256``
so C128 storage is 2. Returning ``[64]`` lets ``select_common_block_size``
split each manager block into four kernel pages. SM100 stays at 256.
"""
if current_platform.is_device_capability_family(120):
return [64]
return [256]


class DeepseekV4SparseMLABackend(AttentionBackend):
"""DeepSeek-V4 sparse-MLA backend base.

Expand All @@ -58,7 +71,7 @@ class DeepseekV4SparseMLABackend(AttentionBackend):

@staticmethod
def get_supported_kernel_block_sizes() -> list[int | MultipleOf]:
return [256]
return dsv4_supported_kernel_block_sizes()

@staticmethod
def get_builder_cls() -> type["DeepseekV4SparseMLAMetadataBuilder"]:
Expand Down
10 changes: 9 additions & 1 deletion vllm/v1/attention/backends/mla/indexer.py
Original file line number Diff line number Diff line change
Expand Up @@ -251,7 +251,15 @@ def supported_kv_cache_layouts(cls) -> tuple[KVCacheLayout, ...]:
@staticmethod
def get_supported_kernel_block_sizes() -> list[int | MultipleOf]:
# Block sizes count uncompressed tokens: C4 indexer pages hold 64 rows.
return [256]
# Imported lazily: vllm.models.deepseek_v4 transitively imports
# vllm._aiter_ops (via fused_moe), which would deadlock cold start
# with an import cycle when vllm._aiter_ops pulls in this module
# first (see vllm/v1/attention/ops/rocm_aiter_mla_sparse.py).
from vllm.models.deepseek_v4.sparse_mla import (
dsv4_supported_kernel_block_sizes,
)

return dsv4_supported_kernel_block_sizes()


class DeepseekV41IndexerBackend(DeepseekV4IndexerBackend):
Expand Down
Loading