@@ -116,41 +116,15 @@ def _run_test(self, kv_seqlens_list, q_seqlens_list, ratio, head_dim, num_states
116116 @pytest .mark .parametrize ('overlap' , [True ], indirect = True )
117117 @pytest .mark .parametrize ('ratio' , [4 ], indirect = True )
118118 @pytest .mark .parametrize ('kv_seqlens_list, q_seqlens_list' , [
119- # decode: history + 1
120- ([13 , 17 , 9 ], [1 , 1 , 1 ]),
121- # decode: various history lengths
122- ([5 , 21 , 33 ], [1 , 1 , 1 ]),
123- # prefill: kv_seqlens == q_seqlens (no history)
124- ([8 , 16 ], [8 , 16 ]),
125- # prefill: long (more than max_write)
126- ([32 , 64 ], [32 , 64 ]),
127- # prefill: with history (multi-turn, start_pos > 0)
128- ([20 , 48 ], [8 , 16 ]),
129- ([128 , 256 ], [16 , 32 ]),
119+ ([13 , 17 , 9 ], [1 , 1 , 1 ]), # decode
120+ ([8 , 16 ], [8 , 16 ]), # prefill: no history
121+ ([20 , 48 ], [8 , 16 ]), # prefill: with history
130122 ])
131123 def test_overlap (self , kv_seqlens_list , q_seqlens_list , ratio , head_dim ,
132124 num_states , overlap , device , dtype ):
133125 self ._run_test (kv_seqlens_list , q_seqlens_list , ratio , head_dim ,
134126 num_states , overlap , device , dtype )
135127
136- # ---- overlap=False, ratio=4 ----
137-
138- @pytest .mark .parametrize ('overlap' , [False ], indirect = True )
139- @pytest .mark .parametrize ('ratio' , [4 ], indirect = True )
140- @pytest .mark .parametrize ('kv_seqlens_list, q_seqlens_list' , [
141- # decode
142- ([13 , 9 , 128 ], [1 , 1 , 1 ]),
143- # prefill
144- ([4 , 8 ], [4 , 8 ]),
145- ([16 ], [16 ]),
146- # prefill: with history
147- ([20 , 40 ], [8 , 16 ]),
148- ])
149- def test_no_overlap (self , kv_seqlens_list , q_seqlens_list , ratio , head_dim ,
150- num_states , overlap , device , dtype ):
151- self ._run_test (kv_seqlens_list , q_seqlens_list , ratio , head_dim ,
152- num_states , overlap , device , dtype )
153-
154128 # ---- overlap=False, ratio=128 (r128 compress path) ----
155129
156130 @pytest .mark .parametrize ('overlap' , [False ], indirect = True )
@@ -439,12 +413,8 @@ def _run_decode_test(self, kvlen, ratio, head_dim, num_states, overlap, device,
439413 @pytest .mark .parametrize ('overlap' , [True ], indirect = True )
440414 @pytest .mark .parametrize ('ratio' , [4 ], indirect = True )
441415 @pytest .mark .parametrize ('kv_seqlens_list, q_seqlens_list' , [
442- # prefill: no history (start_pos=0)
443- ([8 , 16 ], [8 , 16 ]),
444- ([4 , 12 ], [4 , 12 ]),
445- # prefill: with history (start_pos>0)
446- ([20 , 48 ], [8 , 16 ]),
447- ([12 , 24 ], [8 , 16 ]),
416+ ([8 , 16 ], [8 , 16 ]), # no history
417+ ([12 , 24 ], [8 , 16 ]), # with history
448418 ])
449419 def test_prefill_overlap (self , kv_seqlens_list , q_seqlens_list , ratio , head_dim ,
450420 num_states , overlap , device , dtype ):
@@ -456,43 +426,13 @@ def test_prefill_overlap(self, kv_seqlens_list, q_seqlens_list, ratio, head_dim,
456426 @pytest .mark .parametrize ('overlap' , [True ], indirect = True )
457427 @pytest .mark .parametrize ('ratio' , [4 ], indirect = True )
458428 @pytest .mark .parametrize ('kvlen' , [
459- 4 , # first emit (start_pos=3)
460- 8 , # second emit (start_pos=7)
461- 5 , # no emit (start_pos=4, (4+1)%4=1!=0)
462- 7 , # no emit
463- 12 , # third emit
429+ 4 , # emit (start_pos=3)
430+ 5 , # no emit (start_pos=4)
464431 ])
465432 def test_decode_overlap (self , kvlen , ratio , head_dim , num_states , overlap ,
466433 device , dtype ):
467434 self ._run_decode_test (kvlen , ratio , head_dim , num_states , overlap , device , dtype )
468435
469- # ---- overlap=False, ratio=4, prefill ----
470-
471- @pytest .mark .parametrize ('overlap' , [False ], indirect = True )
472- @pytest .mark .parametrize ('ratio' , [4 ], indirect = True )
473- @pytest .mark .parametrize ('kv_seqlens_list, q_seqlens_list' , [
474- ([4 , 8 ], [4 , 8 ]),
475- ([16 ], [16 ]),
476- ([20 , 40 ], [8 , 16 ]),
477- ])
478- def test_prefill_no_overlap (self , kv_seqlens_list , q_seqlens_list , ratio , head_dim ,
479- num_states , overlap , device , dtype ):
480- self ._run_prefill_test (kv_seqlens_list , q_seqlens_list , ratio , head_dim ,
481- num_states , overlap , device , dtype )
482-
483- # ---- overlap=False, ratio=4, decode ----
484-
485- @pytest .mark .parametrize ('overlap' , [False ], indirect = True )
486- @pytest .mark .parametrize ('ratio' , [4 ], indirect = True )
487- @pytest .mark .parametrize ('kvlen' , [
488- 4 , # emit
489- 8 , # emit
490- 5 , # no emit
491- ])
492- def test_decode_no_overlap (self , kvlen , ratio , head_dim , num_states , overlap ,
493- device , dtype ):
494- self ._run_decode_test (kvlen , ratio , head_dim , num_states , overlap , device , dtype )
495-
496436 # ---- overlap=False, ratio=128, prefill ----
497437
498438 @pytest .mark .parametrize ('overlap' , [False ], indirect = True )
@@ -520,6 +460,153 @@ def test_decode_ratio_128(self, kvlen, ratio, head_dim, num_states, overlap,
520460 self ._run_decode_test (kvlen , ratio , head_dim , num_states , overlap , device , dtype )
521461
522462
463+ class TestScoreKVLargeHeadDim :
464+ """Test score_kv with head_dim=512 (n_tiles=4), matching real V4 model
465+ config.
466+
467+ The original tests all use head_dim=128 where n_tiles=1, so a double-d_off bug in the kernel was masked (d_off=0 for
468+ the only tile). Only overlap=True ratio=4 is tested because that is the V4 model's actual config — overlap=False
469+ ratio=4 is not a real config, and overlap=False ratio=128 prefill had no bug (offs_d without d_off prefix).
470+ """
471+
472+ @pytest .fixture
473+ def head_dim (self ):
474+ yield 512
475+
476+ @pytest .fixture
477+ def num_states (self ):
478+ yield 8
479+
480+ @pytest .fixture
481+ def device (self ):
482+ yield 'cuda'
483+
484+ @pytest .fixture
485+ def dtype (self ):
486+ yield torch .bfloat16
487+
488+ @pytest .fixture
489+ def overlap (self , request ):
490+ yield request .param
491+
492+ @pytest .fixture
493+ def ratio (self , request ):
494+ yield request .param
495+
496+ def _run_prefill_test (self , kv_seqlens_list , q_seqlens_list , ratio , head_dim ,
497+ num_states , overlap , device , dtype ):
498+ B = len (kv_seqlens_list )
499+ coff = 1 + overlap
500+ D = coff * head_dim
501+ max_write = ratio * coff
502+ total_q = sum (q_seqlens_list )
503+ max_seqlen_q = max (q_seqlens_list )
504+
505+ kv_seqlens = torch .tensor (kv_seqlens_list , dtype = torch .int32 , device = device )
506+ cu_q_seqlens = torch .tensor ([0 ] + list (q_seqlens_list ), dtype = torch .int32 ,
507+ device = device ).cumsum (0 )
508+
509+ kv = torch .randn (total_q , D , dtype = dtype , device = device )
510+ score = torch .randn (total_q , D , dtype = dtype , device = device )
511+ ape = torch .randn (ratio , D , dtype = torch .float32 , device = device )
512+
513+ state_ids = torch .arange (B , dtype = torch .int32 , device = device )
514+ kv_state = torch .zeros (num_states , max_write , D , dtype = torch .float32 , device = device )
515+ score_state = torch .full ((num_states , max_write , D ), float ('-inf' ),
516+ dtype = torch .float32 , device = device )
517+
518+ from lmdeploy .pytorch .kernels .cuda .v4_compressor import fill_compress_state
519+ has_history = any (kv_seqlens_list [b ] > q_seqlens_list [b ] for b in range (B ))
520+ if has_history :
521+ for b in range (B ):
522+ history_len = kv_seqlens_list [b ] - q_seqlens_list [b ]
523+ if history_len > 0 :
524+ hist_kv = torch .randn (history_len , D , dtype = dtype , device = device )
525+ hist_score = torch .randn (history_len , D , dtype = dtype , device = device )
526+ hist_cu_q = torch .tensor ([0 , history_len ], dtype = torch .int32 , device = device )
527+ hist_kvlen = torch .tensor ([kv_seqlens_list [b ] - q_seqlens_list [b ]],
528+ dtype = torch .int32 , device = device )
529+ hist_sids = torch .tensor ([b ], dtype = torch .int32 , device = device )
530+ fill_compress_state (hist_kv , hist_score , ape , kv_state , score_state ,
531+ hist_sids , hist_cu_q , hist_kvlen )
532+
533+ kv_state_ref = kv_state .clone ()
534+ score_state_ref = score_state .clone ()
535+ ref_compressed = _reference_score_kv (kv .clone (), score .clone (), ape , kv_state_ref ,
536+ score_state_ref , state_ids , cu_q_seqlens ,
537+ kv_seqlens , overlap )
538+
539+ from lmdeploy .pytorch .kernels .cuda .v4_compressor import score_kv
540+ compressed_kv = torch .zeros (total_q , head_dim , dtype = dtype , device = device )
541+ kv_state_k = kv_state .clone ()
542+ score_state_k = score_state .clone ()
543+ score_kv (kv , score , ape , kv_state_k , score_state_k , state_ids ,
544+ cu_q_seqlens , kv_seqlens , compressed_kv , overlap , max_seqlen_q )
545+
546+ torch .testing .assert_close (compressed_kv .float (),
547+ ref_compressed .float (),
548+ atol = 1e-2 , rtol = 1e-2 )
549+
550+ def _run_decode_test (self , kvlen , ratio , head_dim , num_states , overlap , device , dtype ):
551+ coff = 1 + overlap
552+ D = coff * head_dim
553+ max_write = ratio * coff
554+
555+ full_kv = torch .randn (kvlen , D , dtype = dtype , device = device )
556+ full_score = torch .randn (kvlen , D , dtype = dtype , device = device )
557+ ape = torch .randn (ratio , D , dtype = torch .float32 , device = device )
558+
559+ state_ids = torch .tensor ([0 ], dtype = torch .int32 , device = device )
560+ kv_state = torch .zeros (num_states , max_write , D , dtype = torch .float32 , device = device )
561+ score_state = torch .full ((num_states , max_write , D ), float ('-inf' ),
562+ dtype = torch .float32 , device = device )
563+
564+ from lmdeploy .pytorch .kernels .cuda .v4_compressor import fill_compress_state
565+ full_cu_q = torch .tensor ([0 , kvlen ], dtype = torch .int32 , device = device )
566+ full_kvlen = torch .tensor ([kvlen ], dtype = torch .int32 , device = device )
567+ fill_compress_state (full_kv , full_score , ape , kv_state , score_state ,
568+ state_ids , full_cu_q , full_kvlen )
569+
570+ last_kv = full_kv [- 1 :].clone ()
571+ last_score = full_score [- 1 :].clone ()
572+ last_cu_q = torch .tensor ([0 , 1 ], dtype = torch .int32 , device = device )
573+ last_kvlen = torch .tensor ([kvlen ], dtype = torch .int32 , device = device )
574+
575+ kv_state_ref = kv_state .clone ()
576+ score_state_ref = score_state .clone ()
577+ ref_compressed = _reference_score_kv (last_kv .clone (), last_score .clone (), ape ,
578+ kv_state_ref , score_state_ref , state_ids ,
579+ last_cu_q , last_kvlen , overlap )
580+
581+ from lmdeploy .pytorch .kernels .cuda .v4_compressor import score_kv
582+ compressed_kv = torch .zeros (1 , head_dim , dtype = dtype , device = device )
583+ kv_state_k = kv_state .clone ()
584+ score_state_k = score_state .clone ()
585+ score_kv (last_kv , last_score , ape , kv_state_k , score_state_k , state_ids ,
586+ last_cu_q , last_kvlen , compressed_kv , overlap , 1 )
587+
588+ torch .testing .assert_close (compressed_kv .float (),
589+ ref_compressed .float (),
590+ atol = 1e-2 , rtol = 1e-2 )
591+
592+ @pytest .mark .parametrize ('overlap' , [True ], indirect = True )
593+ @pytest .mark .parametrize ('ratio' , [4 ], indirect = True )
594+ @pytest .mark .parametrize ('kvlen' , [12 ])
595+ def test_decode_overlap (self , kvlen , ratio , head_dim , num_states , overlap ,
596+ device , dtype ):
597+ self ._run_decode_test (kvlen , ratio , head_dim , num_states , overlap , device , dtype )
598+
599+ @pytest .mark .parametrize ('overlap' , [True ], indirect = True )
600+ @pytest .mark .parametrize ('ratio' , [4 ], indirect = True )
601+ @pytest .mark .parametrize ('kv_seqlens_list, q_seqlens_list' , [
602+ ([12 , 24 ], [8 , 16 ]),
603+ ])
604+ def test_prefill_overlap (self , kv_seqlens_list , q_seqlens_list , ratio , head_dim ,
605+ num_states , overlap , device , dtype ):
606+ self ._run_prefill_test (kv_seqlens_list , q_seqlens_list , ratio , head_dim ,
607+ num_states , overlap , device , dtype )
608+
609+
523610def _reference_fill_compressed_kv (compressed_kv , kv_cache , cu_q_seqlens , kv_seqlens ,
524611 block_offsets , compress_ratio , block_size ):
525612 """Python reference for fill_compressed_kv, matching
@@ -710,9 +797,8 @@ def _run_decode_test(self, kvlen, compress_ratio, head_dim, block_size, device,
710797
711798 @pytest .mark .parametrize ('compress_ratio' , [4 ], indirect = True )
712799 @pytest .mark .parametrize ('kv_seqlens_list, q_seqlens_list' , [
713- ([8 , 16 ], [8 , 16 ]),
714- ([4 , 12 ], [4 , 12 ]),
715- ([20 , 48 ], [8 , 16 ]),
800+ ([8 , 16 ], [8 , 16 ]), # no history
801+ ([20 , 48 ], [8 , 16 ]), # with history
716802 ])
717803 def test_prefill_r4 (self , kv_seqlens_list , q_seqlens_list , compress_ratio ,
718804 head_dim , block_size , device , dtype ):
@@ -722,7 +808,7 @@ def test_prefill_r4(self, kv_seqlens_list, q_seqlens_list, compress_ratio,
722808 # ---- ratio=4, decode ----
723809
724810 @pytest .mark .parametrize ('compress_ratio' , [4 ], indirect = True )
725- @pytest .mark .parametrize ('kvlen' , [4 , 8 , 5 , 7 , 12 ])
811+ @pytest .mark .parametrize ('kvlen' , [4 , 5 ])
726812 def test_decode_r4 (self , kvlen , compress_ratio , head_dim , block_size , device , dtype ):
727813 self ._run_decode_test (kvlen , compress_ratio , head_dim , block_size , device , dtype )
728814
0 commit comments