|
14 | 14 | import torch |
15 | 15 | import zmq |
16 | 16 | from vllm.utils.network_utils import make_zmq_path |
| 17 | +from vllm.v1.kv_cache_interface import FullAttentionSpec, UniformTypeKVCacheSpecs |
17 | 18 |
|
18 | 19 | fake_engine = types.ModuleType("mooncake.engine") |
19 | 20 | fake_engine.TransferEngine = MagicMock() # type: ignore[attr-defined] |
@@ -295,6 +296,145 @@ def test_reformat_kv_cache_hybrid_linear_uses_cache_block_size(self): |
295 | 296 | torch.testing.assert_close(reformatted_v_cache, expected) |
296 | 297 |
|
297 | 298 |
|
| 299 | +class TestMooncakeTransferGroups(unittest.TestCase): |
| 300 | + def test_build_kv_group2layeridx_splits_uniform_group_by_kv_heads(self): |
| 301 | + mla_spec = FullAttentionSpec( |
| 302 | + block_size=16, |
| 303 | + num_kv_heads=1, |
| 304 | + head_size=64, |
| 305 | + head_size_v=64, |
| 306 | + dtype=torch.float16, |
| 307 | + ) |
| 308 | + qga_spec = FullAttentionSpec( |
| 309 | + block_size=16, |
| 310 | + num_kv_heads=8, |
| 311 | + head_size=64, |
| 312 | + head_size_v=64, |
| 313 | + dtype=torch.float16, |
| 314 | + ) |
| 315 | + layer_specs = { |
| 316 | + "model.layers.0.self_attn": mla_spec, |
| 317 | + "model.layers.1.self_attn": qga_spec, |
| 318 | + } |
| 319 | + uniform_spec = UniformTypeKVCacheSpecs.from_specs(layer_specs) |
| 320 | + assert uniform_spec is not None |
| 321 | + |
| 322 | + worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker) |
| 323 | + worker.vllm_config = MockVllmConfig() |
| 324 | + worker.total_layers = 32 |
| 325 | + worker.kv_cache_config = MockKVCacheConfig( |
| 326 | + kv_cache_groups=[ |
| 327 | + MockKVCacheGroup( |
| 328 | + layer_names=list(layer_specs), |
| 329 | + kv_cache_spec=uniform_spec, |
| 330 | + ) |
| 331 | + ] |
| 332 | + ) |
| 333 | + |
| 334 | + kv_group2layeridx = worker._build_kv_group2layeridx() |
| 335 | + |
| 336 | + self.assertEqual(len(kv_group2layeridx), 2) |
| 337 | + self.assertEqual(kv_group2layeridx[0][0]["kv_cache_group_id"], 0) |
| 338 | + self.assertEqual(kv_group2layeridx[1][0]["kv_cache_group_id"], 0) |
| 339 | + self.assertEqual(worker._get_attention_group_num_key_value_heads(kv_group2layeridx[0][0]), 1) |
| 340 | + self.assertEqual(worker._get_attention_group_num_key_value_heads(kv_group2layeridx[1][0]), 8) |
| 341 | + |
| 342 | + def test_hybrid_rank_pulls_use_transfer_group_kv_heads(self): |
| 343 | + worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker) |
| 344 | + worker.vllm_config = MockVllmConfig() |
| 345 | + worker.vllm_config.model_config.is_deepseek_mla = True |
| 346 | + worker.tp_rank = 0 |
| 347 | + worker.tp_size = 4 |
| 348 | + worker._decode_tp_size = 4 |
| 349 | + worker._prefill_tp_size = 8 |
| 350 | + worker._prefill_pp_size = 1 |
| 351 | + worker.num_key_value_heads = 128 |
| 352 | + worker.use_sparse = False |
| 353 | + worker.kv_group2layeridx = { |
| 354 | + 0: ( |
| 355 | + { |
| 356 | + "kv_cache_spec_type": "UniformTypeKVCacheSpecs", |
| 357 | + "kv_cache_group_id": 0, |
| 358 | + "kv_cache_spec": {"model.layers.0.self_attn": {"num_kv_heads": 1}}, |
| 359 | + }, |
| 360 | + [0], |
| 361 | + ), |
| 362 | + 1: ( |
| 363 | + { |
| 364 | + "kv_cache_spec_type": "UniformTypeKVCacheSpecs", |
| 365 | + "kv_cache_group_id": 0, |
| 366 | + "kv_cache_spec": {"model.layers.1.self_attn": {"num_kv_heads": 8}}, |
| 367 | + }, |
| 368 | + [1], |
| 369 | + ), |
| 370 | + } |
| 371 | + |
| 372 | + _, rank_group_pulls = worker._get_hybrid_remote_rank_group_pulls("req-1", prefill_tp_size=8) |
| 373 | + pulls = [pull for group_pulls in rank_group_pulls.values() for pull in group_pulls] |
| 374 | + mla_pulls = [pull for pull in pulls if pull.group_id == 0] |
| 375 | + qga_pulls = [pull for pull in pulls if pull.group_id == 1] |
| 376 | + |
| 377 | + self.assertEqual(len(mla_pulls), 1) |
| 378 | + self.assertEqual(mla_pulls[0].num_group_pulls, 1) |
| 379 | + self.assertEqual(len(qga_pulls), 2) |
| 380 | + self.assertTrue(all(pull.num_group_pulls == 2 for pull in qga_pulls)) |
| 381 | + |
| 382 | + def test_hybrid_group_pulls_metadata_filters_groups_per_remote_card(self): |
| 383 | + worker = MooncakeConnectorWorker.__new__(MooncakeConnectorWorker) |
| 384 | + worker.vllm_config = MockVllmConfig() |
| 385 | + worker.vllm_config.model_config.is_deepseek_mla = True |
| 386 | + worker._is_hma_required = True |
| 387 | + worker.tp_rank = 0 |
| 388 | + worker.tp_size = 4 |
| 389 | + worker._decode_tp_size = 4 |
| 390 | + worker._prefill_tp_size = 8 |
| 391 | + worker._prefill_pp_size = 1 |
| 392 | + worker.num_key_value_heads = 128 |
| 393 | + worker.use_sparse = False |
| 394 | + worker.kv_group2layeridx = { |
| 395 | + 0: ( |
| 396 | + { |
| 397 | + "kv_cache_spec_type": "FullAttentionSpec", |
| 398 | + "kv_cache_group_id": 0, |
| 399 | + "kv_cache_spec": {"num_kv_heads": 1}, |
| 400 | + }, |
| 401 | + [0], |
| 402 | + ), |
| 403 | + 1: ( |
| 404 | + { |
| 405 | + "kv_cache_spec_type": "FullAttentionSpec", |
| 406 | + "kv_cache_group_id": 0, |
| 407 | + "kv_cache_spec": {"num_kv_heads": 8}, |
| 408 | + }, |
| 409 | + [1], |
| 410 | + ), |
| 411 | + } |
| 412 | + req_id = "req-1" |
| 413 | + remote_base_port = 30000 |
| 414 | + remote_handshake_port_list = [[remote_base_port + rank] for rank in range(8)] |
| 415 | + |
| 416 | + _, expected_rank_group_pulls = worker._get_hybrid_remote_rank_group_pulls(req_id, prefill_tp_size=8) |
| 417 | + group_pulls_list = worker._get_group_pulls_metadata( |
| 418 | + req_id, |
| 419 | + remote_handshake_port_list, |
| 420 | + prefill_tp_size=8, |
| 421 | + remote_base_port=remote_base_port, |
| 422 | + ) |
| 423 | + |
| 424 | + group_ids_by_rank = [ |
| 425 | + [group_pull.group_id for group_pull in group_pulls_list[rank][0]] |
| 426 | + for rank in range(len(remote_handshake_port_list)) |
| 427 | + ] |
| 428 | + expected_group_ids_by_rank = [ |
| 429 | + [group_pull.group_id for group_pull in expected_rank_group_pulls.get(rank, [])] |
| 430 | + for rank in range(len(remote_handshake_port_list)) |
| 431 | + ] |
| 432 | + |
| 433 | + self.assertEqual(group_ids_by_rank, expected_group_ids_by_rank) |
| 434 | + self.assertTrue(any(set(group_ids) != {0, 1} for group_ids in group_ids_by_rank)) |
| 435 | + self.assertFalse(all(set(group_ids) == {0, 1} for group_ids in group_ids_by_rank)) |
| 436 | + |
| 437 | + |
298 | 438 | class TestKVCacheRecvingThreadBasic(unittest.TestCase): |
299 | 439 | def setUp(self): |
300 | 440 | self.engine = MagicMock() |
|
0 commit comments