Skip to content

Commit 887d3de

Browse files
authored
[Bugfix] reject mooncakeV1 c8 consumer KVcache for GQA backend (vllm-project#10947)
### What this PR does / why we need it? Reject MooncakeConnector startup when C8 KV cache quantization is enabled on a KV consumer node with GQA attention backend, because the producer keeps bf16 KV cache while the consumer allocates int8 cache. Add a regression test for the unsupported configuration. ### Does this PR introduce _any_ user-facing change? No ### How was this patch tested? UT test - vLLM version: v0.23.0 - vLLM main: vllm-project/vllm@a30addc --------- Signed-off-by: zzzzzmeng <810924837@qq.com>
1 parent ee8bbbd commit 887d3de

3 files changed

Lines changed: 187 additions & 1 deletion

File tree

tests/ut/test_ascend_config.py

Lines changed: 113 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,9 +15,10 @@
1515

1616
import json
1717
import os
18+
from types import SimpleNamespace
1819
from unittest.mock import patch
1920

20-
from vllm.config import VllmConfig
21+
from vllm.config import KVTransferConfig, VllmConfig
2122

2223
from tests.ut.base import TestBase
2324
from vllm_ascend.ascend_config import clear_ascend_config, get_ascend_config, init_ascend_config
@@ -38,6 +39,20 @@ def wrapper(*args, **kwargs):
3839

3940
return wrapper
4041

42+
@staticmethod
43+
def _make_model_config(
44+
total_num_attention_heads: int = 32,
45+
total_num_kv_heads: int = 8,
46+
is_deepseek_mla: bool = False,
47+
):
48+
return SimpleNamespace(
49+
is_deepseek_mla=is_deepseek_mla,
50+
use_mla=is_deepseek_mla,
51+
enforce_eager=True,
52+
model_arch_config=SimpleNamespace(total_num_attention_heads=total_num_attention_heads),
53+
get_total_num_kv_heads=lambda: total_num_kv_heads,
54+
)
55+
4156
@_clean_up_ascend_config
4257
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
4358
def test_init_ascend_config_without_additional_config(self, mock_fix_incompatible_config):
@@ -94,6 +109,103 @@ def test_init_ascend_config_enable_npugraph_ex(self, mock_fix_incompatible_confi
94109
self.assertTrue(ascend_compilation_config.enable_npugraph_ex)
95110
self.assertTrue(ascend_compilation_config.enable_static_kernel)
96111

112+
@_clean_up_ascend_config
113+
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
114+
def test_init_ascend_config_rejects_mooncake_c8_kv_cache_consumer(self, mock_fix_incompatible_config):
115+
test_vllm_config = VllmConfig()
116+
test_vllm_config.kv_transfer_config = KVTransferConfig(
117+
kv_connector="MooncakeConnectorV1",
118+
kv_role="kv_consumer",
119+
)
120+
test_vllm_config.quant_config = SimpleNamespace(enable_c8_quant=True)
121+
test_vllm_config.model_config = self._make_model_config()
122+
123+
with self.assertRaisesRegex(ValueError, "does not support C8 KV cache quantization"):
124+
init_ascend_config(test_vllm_config)
125+
126+
@_clean_up_ascend_config
127+
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
128+
def test_init_ascend_config_rejects_multi_connector_mooncake_c8_consumer(self, mock_fix_incompatible_config):
129+
test_vllm_config = VllmConfig()
130+
test_vllm_config.kv_transfer_config = KVTransferConfig(
131+
kv_connector="MultiConnector",
132+
kv_role="kv_consumer",
133+
kv_connector_extra_config={
134+
"connectors": [
135+
{
136+
"kv_connector": "MooncakeConnectorV1",
137+
"kv_role": "kv_consumer",
138+
}
139+
]
140+
},
141+
)
142+
test_vllm_config.quant_config = SimpleNamespace(enable_c8_quant=True)
143+
test_vllm_config.model_config = self._make_model_config()
144+
145+
with self.assertRaisesRegex(ValueError, "does not support C8 KV cache quantization"):
146+
init_ascend_config(test_vllm_config)
147+
148+
@_clean_up_ascend_config
149+
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
150+
def test_init_ascend_config_allows_layerwise_c8_kv_cache_consumer(self, mock_fix_incompatible_config):
151+
test_vllm_config = VllmConfig()
152+
test_vllm_config.kv_transfer_config = KVTransferConfig(
153+
kv_connector="MooncakeLayerwiseConnector",
154+
kv_role="kv_consumer",
155+
)
156+
test_vllm_config.quant_config = SimpleNamespace(enable_c8_quant=True)
157+
test_vllm_config.model_config = self._make_model_config()
158+
159+
ascend_config = init_ascend_config(test_vllm_config)
160+
161+
self.assertIsNotNone(ascend_config)
162+
163+
@_clean_up_ascend_config
164+
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
165+
def test_init_ascend_config_allows_mha_mooncake_c8_kv_cache_consumer(self, mock_fix_incompatible_config):
166+
test_vllm_config = VllmConfig()
167+
test_vllm_config.kv_transfer_config = KVTransferConfig(
168+
kv_connector="MooncakeConnectorV1",
169+
kv_role="kv_consumer",
170+
)
171+
test_vllm_config.quant_config = SimpleNamespace(enable_c8_quant=True)
172+
test_vllm_config.model_config = self._make_model_config(
173+
total_num_attention_heads=8,
174+
total_num_kv_heads=8,
175+
)
176+
177+
ascend_config = init_ascend_config(test_vllm_config)
178+
179+
self.assertIsNotNone(ascend_config)
180+
181+
@_clean_up_ascend_config
182+
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
183+
def test_init_ascend_config_rejects_mooncake_c8_kv_cache_producer(self, mock_fix_incompatible_config):
184+
test_vllm_config = VllmConfig()
185+
test_vllm_config.kv_transfer_config = KVTransferConfig(
186+
kv_connector="MooncakeConnectorV1",
187+
kv_role="kv_producer",
188+
)
189+
test_vllm_config.quant_config = SimpleNamespace(enable_c8_quant=True)
190+
test_vllm_config.model_config = self._make_model_config()
191+
192+
with self.assertRaisesRegex(ValueError, "does not support C8 KV cache quantization"):
193+
init_ascend_config(test_vllm_config)
194+
195+
@_clean_up_ascend_config
196+
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")
197+
def test_init_ascend_config_rejects_mooncake_c8_kv_cache_both_role(self, mock_fix_incompatible_config):
198+
test_vllm_config = VllmConfig()
199+
test_vllm_config.kv_transfer_config = KVTransferConfig(
200+
kv_connector="MooncakeConnectorV1",
201+
kv_role="kv_both",
202+
)
203+
test_vllm_config.quant_config = SimpleNamespace(enable_c8_quant=True)
204+
test_vllm_config.model_config = self._make_model_config()
205+
206+
with self.assertRaisesRegex(ValueError, "does not support C8 KV cache quantization"):
207+
init_ascend_config(test_vllm_config)
208+
97209
@_clean_up_ascend_config
98210
@patch("vllm_ascend.ascend_config.logger.info_once")
99211
@patch("vllm_ascend.platform.NPUPlatform.check_and_update_config")

vllm_ascend/ascend_config.py

Lines changed: 27 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,6 +32,7 @@ class AscendConfig:
3232
def __init__(self, vllm_config: "VllmConfig"):
3333
self.vllm_config = vllm_config
3434
additional_config = vllm_config.additional_config if vllm_config.additional_config is not None else {}
35+
self._check_mooncake_c8_kv_cache_quant(vllm_config)
3536

3637
xlite_graph_config = additional_config.get("xlite_graph_config", {})
3738
self.xlite_graph_config = XliteGraphConfig(xlite_graph_config, vllm_config)
@@ -320,6 +321,32 @@ def _get_config_value(additional_config: dict[str, Any], config_key: str, env_ke
320321
)
321322
return env_value
322323

324+
@classmethod
325+
def _check_mooncake_c8_kv_cache_quant(cls, vllm_config: "VllmConfig") -> None:
326+
kv_transfer_config = getattr(vllm_config, "kv_transfer_config", None)
327+
if kv_transfer_config is None:
328+
return
329+
330+
quant_config = getattr(vllm_config, "quant_config", None)
331+
enable_c8_quant = getattr(quant_config, "enable_c8_quant", False)
332+
if enable_c8_quant is not True:
333+
return
334+
335+
from vllm_ascend.utils import is_gqa_backend, uses_mooncake_connector
336+
337+
if not is_gqa_backend(vllm_config):
338+
return
339+
340+
if not uses_mooncake_connector(kv_transfer_config):
341+
return
342+
343+
raise ValueError(
344+
"MooncakeConnector does not support C8 KV cache quantization on GQA models. "
345+
"The producer keeps KV cache in bf16 while the consumer allocates int8 KV cache, so raw "
346+
"Mooncake transfer would reinterpret bf16 bytes as int8. Please disable C8 KV cache quantization "
347+
"or use MooncakeLayerwiseConnector, which quantizes KV cache before transfer."
348+
)
349+
323350
def _check_mix_placement(self):
324351
if self.mix_placement:
325352
if self.enable_shared_expert_dp or self.multistream_overlap_shared_expert:

vllm_ascend/utils.py

Lines changed: 47 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1404,6 +1404,53 @@ def _check(name: str, config: dict):
14041404
_check("decode", vllm_config.kv_transfer_config.get_from_extra_config("decode", {}))
14051405

14061406

1407+
def is_gqa_backend(vllm_config: VllmConfig) -> bool:
1408+
model_config = getattr(vllm_config, "model_config", None)
1409+
if model_config is None:
1410+
return False
1411+
1412+
if getattr(model_config, "is_deepseek_mla", False) or getattr(model_config, "use_mla", False):
1413+
return False
1414+
1415+
model_arch_config = getattr(model_config, "model_arch_config", None)
1416+
total_num_attention_heads = getattr(model_arch_config, "total_num_attention_heads", None)
1417+
get_total_num_kv_heads = getattr(model_config, "get_total_num_kv_heads", None)
1418+
if total_num_attention_heads is None or not callable(get_total_num_kv_heads):
1419+
return False
1420+
1421+
total_num_kv_heads = get_total_num_kv_heads()
1422+
if total_num_kv_heads is None:
1423+
return False
1424+
1425+
return total_num_attention_heads != total_num_kv_heads
1426+
1427+
1428+
def uses_mooncake_connector(kv_transfer_config: Any) -> bool:
1429+
mooncake_connector_names = {"MooncakeConnector", "MooncakeConnectorV1"}
1430+
return bool(_collect_kv_connector_names(kv_transfer_config) & mooncake_connector_names)
1431+
1432+
1433+
def _collect_kv_connector_names(value: Any) -> set[str]:
1434+
connector_names: set[str] = set()
1435+
if isinstance(value, dict):
1436+
connector = value.get("kv_connector")
1437+
if isinstance(connector, str):
1438+
connector_names.add(connector)
1439+
for nested_value in value.values():
1440+
connector_names.update(_collect_kv_connector_names(nested_value))
1441+
elif isinstance(value, (list, tuple)):
1442+
for nested_value in value:
1443+
connector_names.update(_collect_kv_connector_names(nested_value))
1444+
else:
1445+
connector = getattr(value, "kv_connector", None)
1446+
if isinstance(connector, str):
1447+
connector_names.add(connector)
1448+
extra_config = getattr(value, "kv_connector_extra_config", None)
1449+
if isinstance(extra_config, (dict, list, tuple)):
1450+
connector_names.update(_collect_kv_connector_names(extra_config))
1451+
return connector_names
1452+
1453+
14071454
def singleton(cls):
14081455
instances = {}
14091456

0 commit comments

Comments
 (0)