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
44 changes: 44 additions & 0 deletions tests/models/kimi_k3/test_dspark_kv_group.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,44 @@
# SPDX-License-Identifier: Apache-2.0
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
"""The K3 DSpark draft must not merge into the target's MLA KV cache group."""

from types import SimpleNamespace

import pytest

from vllm.models.kimi_k3.nvidia.mla import MultiHeadLatentAttention
from vllm.v1.kv_cache_interface import MLAAttentionSpec


def _get_spec(non_causal_multi_token_decode: bool) -> MLAAttentionSpec:
mla = MultiHeadLatentAttention.__new__(MultiHeadLatentAttention)
mla.kv_cache_dtype = "auto"
mla.head_size = 576
mla.non_causal_multi_token_decode = non_causal_multi_token_decode
vllm_config = SimpleNamespace(
cache_config=SimpleNamespace(block_size=64),
model_config=None,
)
return mla.get_kv_cache_spec(vllm_config) # type: ignore[arg-type]


def test_dspark_draft_spec_carries_model_version():
spec = _get_spec(non_causal_multi_token_decode=True)
assert spec.model_version == "kimi_k3_dspark"


def test_target_spec_has_no_model_version():
spec = _get_spec(non_causal_multi_token_decode=False)
assert spec.model_version is None


def test_dspark_draft_spec_cannot_merge_into_target_group():
# The split matters because MLAAttentionSpec.merge ORs
# non_causal_multi_token_decode, which would flag the causal target.
with pytest.raises(AssertionError, match="model version"):
MLAAttentionSpec.merge(
[
_get_spec(non_causal_multi_token_decode=False),
_get_spec(non_causal_multi_token_decode=True),
]
)
9 changes: 9 additions & 0 deletions vllm/models/kimi_k3/nvidia/mla.py
Original file line number Diff line number Diff line change
Expand Up @@ -430,6 +430,15 @@ def get_kv_cache_spec(self, vllm_config: VllmConfig) -> KVCacheSpec:
kv_quant_mode=get_kv_quant_mode(self.kv_cache_dtype),
# fp8_ds_mla: 656-byte custom layout; see flashmla_sparse.py.
state_content_bytes=656 if self.kv_cache_dtype == "fp8_ds_mla" else None,
# Keep the non-causal DSpark draft out of the target's KV cache
# group: MLAAttentionSpec.merge ORs non_causal_multi_token_decode,
# so a merged group would flag the causal target too, raising its
# TritonMLA reorder threshold and misrouting its short prefills and
# causal verification blocks into a decode path that expects one
# query row per request.
model_version="kimi_k3_dspark"
if self.non_causal_multi_token_decode
else None,
Comment on lines +439 to +441

@GirasoleY GirasoleY Sep 15, 2026

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I feel we may still want to allocate the kv cache as one group, but fix trtion mla behavior.

A quick search in community it feels #51065 is closer to a proper fix. Could you check if this PR fix your use case?

non_causal_multi_token_decode=self.non_causal_multi_token_decode,
)

Expand Down
Loading