2323 MooncakeXferMetadata ,
2424 SendBlockMeta ,
2525 TransferRegion ,
26+ _compute_sender_transfer_plan ,
27+ _validate_asymmetric_region_lengths ,
2628)
2729from vllm .v1 .attention .backends .registry import MambaAttentionBackendEnum
2830from 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