diff --git a/tests/models/kimi_k3/test_dspark_kv_group.py b/tests/models/kimi_k3/test_dspark_kv_group.py new file mode 100644 index 000000000000..bd0299d47b80 --- /dev/null +++ b/tests/models/kimi_k3/test_dspark_kv_group.py @@ -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), + ] + ) diff --git a/vllm/models/kimi_k3/nvidia/mla.py b/vllm/models/kimi_k3/nvidia/mla.py index 87bf59749c1d..4201d890354f 100644 --- a/vllm/models/kimi_k3/nvidia/mla.py +++ b/vllm/models/kimi_k3/nvidia/mla.py @@ -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, non_causal_multi_token_decode=self.non_causal_multi_token_decode, )