Skip to content

Commit 6a2e400

Browse files
committed
[Bugfix][KV Connector] Mooncake: heterogeneous-TP support for hybrid GDN/Mamba models
- Register every mamba state (conv, ssm) as its own region with its real unpadded byte length so TP-ratio slicing stays head-aligned. - Split the GDN conv state into per-sub-projection (Q/K/V) regions under the DS layout, reusing derive_mamba_conv_split; a single conv region let the het-TP byte split cut across projection boundaries and silently scrambled the conv state on consumer ranks (GSM8K -10pp with occasional empty generations on Qwen3.5-397B PP2TP2->TP4). - Copy KV regions whole when the consumer TP group is wider than the KV-head count (consumer ranks replicate head groups, not slice them). - Tolerate peer-PP-stage layers absent from the local spec table. - Group-aware per-region asymmetric length validation. Validated on 2-node Qwen3.5-397B-A17B-NVFP4 PP2TP2->TP4 Mooncake/RDMA: GSM8K strict-match 0.9409 (aggregate TP4 baseline 0.9363).
1 parent f1178f3 commit 6a2e400

2 files changed

Lines changed: 358 additions & 31 deletions

File tree

tests/v1/kv_connector/unit/test_mooncake_connector_hybrid_mamba.py

Lines changed: 253 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -23,6 +23,8 @@
2323
MooncakeXferMetadata,
2424
SendBlockMeta,
2525
TransferRegion,
26+
_compute_sender_transfer_plan,
27+
_validate_asymmetric_region_lengths,
2628
)
2729
from vllm.v1.attention.backends.registry import MambaAttentionBackendEnum
2830
from vllm.v1.attention.backends.utils import NULL_BLOCK_ID
@@ -139,26 +141,33 @@ def test_register_kv_caches_emits_fa_and_gdn_regions(monkeypatch):
139141
worker = connector.connector_worker
140142

141143
fa_cache = torch.empty((2, 2, 11), dtype=torch.float16)
142-
gdn_conv_state = torch.empty((2, 22), dtype=torch.float16)
143-
gdn_ssm_state = torch.empty((2, 4), dtype=torch.float16)
144+
gdn_page_bytes = kv_cache_config.kv_cache_groups[
145+
1
146+
].kv_cache_spec.page_size_bytes
147+
gdn_cache = torch.empty((2, 1, 1, gdn_page_bytes), dtype=torch.int8)
144148

145149
worker.register_kv_caches(
146150
{
147151
"model.layers.0.self_attn": fa_cache,
148-
"model.layers.1.linear_attn": (gdn_conv_state, gdn_ssm_state),
152+
"model.layers.1.linear_attn": gdn_cache,
149153
}
150154
)
151155

152156
assert worker.transfer_topo.is_mamba is True
153157
assert worker.registered_layer_names == [
154158
"model.layers.0.self_attn",
155159
"model.layers.1.linear_attn",
160+
"model.layers.1.linear_attn",
156161
]
157-
assert worker.registered_group_indices == [0, 1]
162+
assert worker.registered_group_indices == [0, 1, 1]
163+
# The GDN page packs conv (36 B) then ssm (8 B); each state registers
164+
# as its own region with real, unpadded byte length.
158165
assert worker.kv_caches_base_addr == [
159166
fa_cache.data_ptr(),
160-
gdn_conv_state.data_ptr(),
167+
gdn_cache.data_ptr(),
168+
gdn_cache.data_ptr() + 36,
161169
]
170+
assert worker.kv_block_len_per_layer == [22, 36, 8]
162171

163172
worker.shutdown()
164173
worker.shutdown = noop_shutdown
@@ -185,22 +194,22 @@ def test_register_kv_caches_deduplicates_shared_backing_memory(monkeypatch):
185194

186195
backing = torch.empty((4, 64), dtype=torch.float16)
187196
fa_cache = backing[:2, :16]
188-
gdn_conv_state = backing[:3]
189-
gdn_ssm_state = torch.empty((3, 4), dtype=torch.float16)
197+
gdn_cache = backing[:3].unsqueeze(1).unsqueeze(1)
190198

191199
with patch.object(
192200
worker.engine, "batch_register_memory", return_value=0
193201
) as batch_register_memory:
194202
worker.register_kv_caches(
195203
{
196204
"model.layers.0.self_attn": fa_cache,
197-
"model.layers.1.linear_attn": (gdn_conv_state, gdn_ssm_state),
205+
"model.layers.1.linear_attn": gdn_cache,
198206
}
199207
)
200208

201209
assert worker.kv_caches_base_addr == [
202210
fa_cache.data_ptr(),
203-
gdn_conv_state.data_ptr(),
211+
gdn_cache.data_ptr(),
212+
gdn_cache.data_ptr() + 36,
204213
]
205214
batch_register_memory.assert_called_once()
206215
registered_ptrs, registered_lens = batch_register_memory.call_args[0]
@@ -384,3 +393,238 @@ def test_hybrid_gdn_splits_fa_regions_but_keeps_gdn_state_whole(
384393
worker.shutdown()
385394
worker.shutdown = noop_shutdown
386395
connector.connector_worker = None
396+
397+
398+
def test_get_transfer_regions_tolerates_peer_pp_stage_layers(monkeypatch):
399+
# A PP-sharded worker's local spec table covers only its own Mamba/GDN
400+
# layers, but remote metadata lists every layer. Expanding remote regions
401+
# for a peer stage's layers must not raise KeyError; the split flag for
402+
# such layers is unused because alignment matches local layers only.
403+
monkeypatch.setenv("VLLM_MOONCAKE_ABORT_REQUEST_TIMEOUT", "5")
404+
vllm_config = create_vllm_config(
405+
kv_connector="MooncakeConnector",
406+
kv_role="kv_producer",
407+
)
408+
kv_cache_config = make_hybrid_gdn_kv_cache_config(
409+
vllm_config.cache_config.block_size
410+
)
411+
412+
with set_current_vllm_config(vllm_config), patch_worker_dependencies():
413+
connector = MooncakeConnector(
414+
vllm_config,
415+
KVConnectorRole.WORKER,
416+
kv_cache_config,
417+
)
418+
worker = connector.connector_worker
419+
420+
worker.transfer_topo = SimpleNamespace(virtually_split_kv_in_blocks=True)
421+
regions = worker._get_transfer_regions(
422+
base_addrs=[0x1000, 0x2000, 0x3000],
423+
block_lens=[0x100, 0x100, 0x100],
424+
kv_block_lens=[0x40, 0x100, 0x100],
425+
layer_names=[
426+
"model.layers.0.self_attn",
427+
"model.layers.1.linear_attn",
428+
"model.layers.31.linear_attn",
429+
],
430+
layer_indices=[0, 1, 31],
431+
group_indices=[0, 1, 1],
432+
)
433+
434+
assert [
435+
(region.group_index, region.base_addr, region.kv_block_len)
436+
for region in regions
437+
] == [
438+
(0, 0x1000, 0x40),
439+
(0, 0x1040, 0x40),
440+
(1, 0x2000, 0x100),
441+
(1, 0x3000, 0x100),
442+
(1, 0x3100, 0x100),
443+
]
444+
445+
worker.shutdown()
446+
worker.shutdown = noop_shutdown
447+
connector.connector_worker = None
448+
449+
450+
def test_transfer_plan_full_copy_for_replicated_consumer_kv():
451+
# P TP2 -> D TP4 with 2 KV heads: every consumer rank holds a full replica
452+
# of its head group's attention region instead of a slice of the
453+
# producer's.
454+
assert _compute_sender_transfer_plan(
455+
local_tp_rank=0,
456+
local_tp_size=2,
457+
remote_tp_rank=3,
458+
remote_tp_size=4,
459+
local_kv_block_len=1000,
460+
remote_kv_block_len=1000,
461+
producer_cache_replicated=False,
462+
consumer_kv_replicated=True,
463+
) == (True, 0, 0, 1000)
464+
# Mamba/GDN states shard by head count and keep ratio slicing.
465+
assert _compute_sender_transfer_plan(
466+
local_tp_rank=0,
467+
local_tp_size=2,
468+
remote_tp_rank=3,
469+
remote_tp_size=4,
470+
local_kv_block_len=2000,
471+
remote_kv_block_len=1000,
472+
producer_cache_replicated=False,
473+
consumer_kv_replicated=False,
474+
) == (True, 1000, 0, 1000)
475+
476+
477+
def test_validate_region_lengths_allows_replicated_consumer_attention():
478+
kv_cache_config = make_hybrid_gdn_kv_cache_config(block_size=16)
479+
local_regions = [
480+
TransferRegion("model.layers.0.self_attn", 0, 0x1000, 2000, 1000, 0),
481+
TransferRegion("model.layers.1.linear_attn", 1, 0x5000, 2000, 2000, 1),
482+
]
483+
remote_regions = [
484+
TransferRegion("model.layers.0.self_attn", 0, 0x2000, 2000, 1000, 0),
485+
TransferRegion("model.layers.1.linear_attn", 1, 0x6000, 2000, 1000, 1),
486+
]
487+
# Replicated consumer KV: attention regions must match exactly, Mamba/GDN
488+
# regions keep the TP-ratio rule (local == 2x remote for TP2 -> TP4).
489+
assert (
490+
_validate_asymmetric_region_lengths(
491+
local_regions,
492+
remote_regions,
493+
local_tp_size=2,
494+
remote_tp_size=4,
495+
producer_cache_replicated=False,
496+
group_specs=kv_cache_config.kv_cache_groups,
497+
total_num_kv_heads=2,
498+
)
499+
is None
500+
)
501+
502+
mismatched_remote = [
503+
TransferRegion("model.layers.0.self_attn", 0, 0x2000, 2000, 500, 0),
504+
remote_regions[1],
505+
]
506+
assert "replicated consumer KV" in _validate_asymmetric_region_lengths(
507+
local_regions,
508+
mismatched_remote,
509+
local_tp_size=2,
510+
remote_tp_size=4,
511+
producer_cache_replicated=False,
512+
group_specs=kv_cache_config.kv_cache_groups,
513+
total_num_kv_heads=2,
514+
)
515+
516+
517+
@pytest.mark.cpu_test
518+
def test_register_kv_caches_splits_gdn_conv_sub_projections(monkeypatch):
519+
"""DS conv layout: each GDN conv sub-projection is its own region.
520+
521+
Het-TP slicing splits every region by TP ratio, so the conv page must be
522+
decomposed into Q/K/V regions; a single conv region would let the ratio
523+
split cut across projection boundaries and scramble the state.
524+
"""
525+
from vllm.model_executor.layers.mamba.mamba_utils import (
526+
get_conv_state_layout,
527+
)
528+
529+
monkeypatch.setenv("VLLM_MOONCAKE_ABORT_REQUEST_TIMEOUT", "5")
530+
monkeypatch.setenv("VLLM_SSM_CONV_STATE_LAYOUT", "DS")
531+
get_conv_state_layout.cache_clear()
532+
try:
533+
vllm_config = create_vllm_config(
534+
kv_connector="MooncakeConnector",
535+
kv_role="kv_consumer",
536+
)
537+
kv_cache_config = make_hybrid_gdn_kv_cache_config(
538+
vllm_config.cache_config.block_size
539+
)
540+
541+
with set_current_vllm_config(vllm_config), patch_worker_dependencies():
542+
connector = MooncakeConnector(
543+
vllm_config,
544+
KVConnectorRole.WORKER,
545+
kv_cache_config,
546+
)
547+
worker = connector.connector_worker
548+
549+
fa_cache = torch.empty((2, 2, 11), dtype=torch.float16)
550+
gdn_page_bytes = kv_cache_config.kv_cache_groups[
551+
1
552+
].kv_cache_spec.page_size_bytes
553+
gdn_cache = torch.empty((2, 1, 1, gdn_page_bytes), dtype=torch.int8)
554+
555+
worker.register_kv_caches(
556+
{
557+
"model.layers.0.self_attn": fa_cache,
558+
"model.layers.1.linear_attn": gdn_cache,
559+
}
560+
)
561+
562+
# Fixture GDN spec: conv (6, 3) fp16 -> Q/K/V dims (2, 2, 2),
563+
# 12 B each in DS layout; ssm (1, 2, 2) fp16 -> 8 B at offset 36.
564+
assert worker.registered_layer_names == [
565+
"model.layers.0.self_attn",
566+
*["model.layers.1.linear_attn"] * 4,
567+
]
568+
assert worker.registered_group_indices == [0, 1, 1, 1, 1]
569+
assert worker.kv_caches_base_addr == [
570+
fa_cache.data_ptr(),
571+
gdn_cache.data_ptr(),
572+
gdn_cache.data_ptr() + 12,
573+
gdn_cache.data_ptr() + 24,
574+
gdn_cache.data_ptr() + 36,
575+
]
576+
assert worker.kv_block_len_per_layer == [22, 12, 12, 12, 8]
577+
578+
worker.shutdown()
579+
worker.shutdown = noop_shutdown
580+
connector.connector_worker = None
581+
finally:
582+
get_conv_state_layout.cache_clear()
583+
584+
585+
def test_gdn_conv_sub_projection_regions_align_across_het_tp(monkeypatch):
586+
"""P TP2 -> D TP4 GDN: per-region ratio slices land inside projections."""
587+
from vllm.distributed.kv_transfer.kv_connector.v1.ssm_conv_transfer_utils import ( # noqa: E501
588+
derive_mamba_conv_split,
589+
)
590+
from vllm.model_executor.layers.mamba.mamba_utils import (
591+
get_conv_state_layout,
592+
)
593+
594+
monkeypatch.setenv("VLLM_SSM_CONV_STATE_LAYOUT", "DS")
595+
get_conv_state_layout.cache_clear()
596+
try:
597+
# Global GDN dims: key_dim=4 (Q == K), value_dim=8 -> conv_dim=16.
598+
def gdn_spec(tp: int) -> MambaSpec:
599+
return MambaSpec(
600+
block_size=16,
601+
shapes=((16 // tp, 3), (8 // tp // 2, 2, 2)),
602+
dtypes=(torch.float16, torch.float16),
603+
mamba_type=MambaAttentionBackendEnum.GDN_ATTN,
604+
)
605+
606+
p_split = derive_mamba_conv_split(gdn_spec(tp=2), local_tp=2)
607+
d_split = derive_mamba_conv_split(gdn_spec(tp=4), local_tp=4)
608+
assert p_split.local_proj_dims == (2, 2, 4)
609+
assert d_split.local_proj_dims == (1, 1, 2)
610+
611+
# Every sub-projection region obeys the TP-ratio rule (P len == 2x D
612+
# len), and each D rank slices its own half of each P sub-projection.
613+
row_bytes = 3 * 2 # conv_rows x fp16
614+
for p_dim, d_dim in zip(p_split.local_proj_dims, d_split.local_proj_dims):
615+
p_len = p_dim * row_bytes
616+
d_len = d_dim * row_bytes
617+
assert p_len == 2 * d_len
618+
for d_rank in range(4):
619+
assert _compute_sender_transfer_plan(
620+
local_tp_rank=d_rank // 2,
621+
local_tp_size=2,
622+
remote_tp_rank=d_rank,
623+
remote_tp_size=4,
624+
local_kv_block_len=p_len,
625+
remote_kv_block_len=d_len,
626+
producer_cache_replicated=False,
627+
consumer_kv_replicated=False,
628+
) == (True, (d_rank % 2) * d_len, 0, d_len)
629+
finally:
630+
get_conv_state_layout.cache_clear()

0 commit comments

Comments
 (0)