@@ -898,25 +898,25 @@ def blocked_kv_mla_attention_forward(
898898 absorption = False
899899
900900 k_heads , q_heads = compressed_kv_block .shape [1 ], query .shape [1 ]
901- num_heads_to_repeat = q_heads - k_heads
902- repeated_ckv_block = compressed_kv_block [:, 0 , :, :].expand (
903- batch_size , num_heads_to_repeat , - 1 , module .kv_lora_rank
904- )
905- compressed_kv_block = torch .cat ((compressed_kv_block , repeated_ckv_block ), dim = 1 )
906901
907- repeated_k_pe_block = k_pe_block [:, 0 , :, :].expand (
908- batch_size , num_heads_to_repeat , - 1 , module .qk_rope_head_dim
909- )
910- k_pe_block = torch .cat ((k_pe_block , repeated_k_pe_block ), dim = 1 )
902+ if k_heads > 1 :
903+ num_heads_to_repeat = math .ceil (q_heads / k_heads )
904+ compressed_kv_block = (
905+ compressed_kv_block .unsqueeze (2 )
906+ .expand (- 1 , - 1 , num_heads_to_repeat , - 1 , - 1 )
907+ .reshape (batch_size , num_heads_to_repeat * k_heads , - 1 , module .config .kv_lora_rank )
908+ )
909+ compressed_kv_block = compressed_kv_block [:, :q_heads , :, :]
910+
911+ k_pe_block = (
912+ k_pe_block .unsqueeze (2 )
913+ .expand (- 1 , - 1 , num_heads_to_repeat , - 1 , - 1 )
914+ .reshape (batch_size , num_heads_to_repeat * k_heads , - 1 , module .config .qk_rope_head_dim )
915+ )
916+ k_pe_block = k_pe_block [:, :q_heads , :, :]
911917
912918 if absorption :
913919 krope_nope = torch .cat ((compressed_kv_block , k_pe_block ), dim = - 1 )
914- k_heads , q_heads = krope_nope .shape [1 ], query .shape [1 ]
915- num_heads_to_repeat = q_heads - k_heads
916- repeated_k = krope_nope [:, 0 , :, :].expand (
917- batch_size , num_heads_to_repeat , - 1 , module .qk_rope_head_dim + module .kv_lora_rank
918- )
919- krope_nope = torch .cat ((krope_nope , repeated_k ), dim = 1 )
920920 attn_weights_block = torch .matmul (query , krope_nope .transpose (2 , 3 )) * scaling
921921 # [1, 64, q_len, 576] X [1, 1, 576, kv_block_size] -> [1, 64, q_len, kv_block_size]
922922 attn_weights_block = torch .where (causal_mask_block , masked_tensor , attn_weights_block )
@@ -930,20 +930,8 @@ def blocked_kv_mla_attention_forward(
930930 skip_future ,
931931 ) # [1, 64, q_len, kv_block_size] X [1, 1, kv_block_size, 512] -> [1, 64, q_len, 512]
932932 else :
933- k_heads , q_heads = compressed_kv_block .shape [1 ], query .shape [1 ]
934- num_heads_to_repeat = q_heads - k_heads
935- repeated_ckv_block = compressed_kv_block [:, 0 , :, :].expand (
936- batch_size , num_heads_to_repeat , - 1 , module .kv_lora_rank
937- )
938- compressed_kv_block = torch .cat ((compressed_kv_block , repeated_ckv_block ), dim = 1 )
939933 knope = torch .matmul (compressed_kv_block , per_head_k_up_normal )
940-
941- repeated_k_pe_block = k_pe_block [:, 0 , :, :].expand (
942- batch_size , num_heads_to_repeat , - 1 , module .qk_rope_head_dim
943- )
944- k_pe_block = torch .cat ((k_pe_block , repeated_k_pe_block ), dim = 1 )
945-
946- krope_nope = torch .cat ((knope , k_pe_block .expand (- 1 , num_heads , - 1 , - 1 )), dim = - 1 )
934+ krope_nope = torch .cat ((knope , k_pe_block ), dim = - 1 )
947935 attn_weights_block = torch .matmul (query , krope_nope .transpose (2 , 3 )) * scaling
948936 attn_weights_block = torch .where (causal_mask_block , masked_tensor , attn_weights_block )
949937 current_max , current_denominator , output = update_running_softmax (
0 commit comments