CUDA: fix Gemma E4B MTP FlashAttention (#25148)
* CUDA: fix Gemma E4B MTP FlashAttention * remove unused template declaration
This commit is contained in:
@@ -2003,6 +2003,10 @@ DECL_FATTN_MMA_F16_CASE_ALL_NCOLS2(112, 112, 64)
|
|||||||
DECL_FATTN_MMA_F16_CASE_ALL_NCOLS2(128, 128, 64)
|
DECL_FATTN_MMA_F16_CASE_ALL_NCOLS2(128, 128, 64)
|
||||||
DECL_FATTN_MMA_F16_CASE_ALL_NCOLS2(256, 256, 64)
|
DECL_FATTN_MMA_F16_CASE_ALL_NCOLS2(256, 256, 64)
|
||||||
|
|
||||||
|
extern DECL_FATTN_MMA_F16_CASE(512, 512, 4, 2);
|
||||||
|
extern DECL_FATTN_MMA_F16_CASE(512, 512, 8, 2);
|
||||||
|
extern DECL_FATTN_MMA_F16_CASE(512, 512, 16, 2);
|
||||||
|
extern DECL_FATTN_MMA_F16_CASE(512, 512, 32, 2);
|
||||||
extern DECL_FATTN_MMA_F16_CASE(512, 512, 2, 4);
|
extern DECL_FATTN_MMA_F16_CASE(512, 512, 2, 4);
|
||||||
extern DECL_FATTN_MMA_F16_CASE(512, 512, 4, 4);
|
extern DECL_FATTN_MMA_F16_CASE(512, 512, 4, 4);
|
||||||
extern DECL_FATTN_MMA_F16_CASE(512, 512, 8, 4);
|
extern DECL_FATTN_MMA_F16_CASE(512, 512, 8, 4);
|
||||||
|
|||||||
@@ -76,6 +76,7 @@ static constexpr __host__ __device__ uint32_t ggml_cuda_fattn_tile_get_config_nv
|
|||||||
|
|
||||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(320, 256, 16, 256, 2, 64, 64)
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(320, 256, 16, 256, 2, 64, 64)
|
||||||
|
|
||||||
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 2, 64, 2, 64, 64)
|
||||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 4, 128, 2, 64, 64)
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 4, 128, 2, 64, 64)
|
||||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 8, 256, 2, 64, 64)
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 8, 256, 2, 64, 64)
|
||||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 16, 256, 2, 64, 64)
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 16, 256, 2, 64, 64)
|
||||||
@@ -144,6 +145,7 @@ static constexpr __host__ __device__ uint32_t ggml_cuda_fattn_tile_get_config_nv
|
|||||||
|
|
||||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(320, 256, 16, 256, 2, 32, 64)
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(320, 256, 16, 256, 2, 32, 64)
|
||||||
|
|
||||||
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 2, 64, 2, 32, 64)
|
||||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 4, 128, 2, 32, 64)
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 4, 128, 2, 32, 64)
|
||||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 8, 256, 2, 32, 64)
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 8, 256, 2, 32, 64)
|
||||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 16, 256, 2, 32, 64)
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 16, 256, 2, 32, 64)
|
||||||
@@ -219,6 +221,7 @@ static constexpr __host__ __device__ uint32_t ggml_cuda_fattn_tile_get_config_am
|
|||||||
|
|
||||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(320, 256, 32, 512, 1, 128, 64)
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(320, 256, 32, 512, 1, 128, 64)
|
||||||
|
|
||||||
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 2, 64, 2, 64, 64)
|
||||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 4, 128, 2, 64, 64)
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 4, 128, 2, 64, 64)
|
||||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 8, 256, 2, 64, 64)
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 8, 256, 2, 64, 64)
|
||||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 16, 256, 2, 64, 64)
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 16, 256, 2, 64, 64)
|
||||||
@@ -296,6 +299,7 @@ static constexpr __host__ __device__ uint32_t ggml_cuda_fattn_tile_get_config_am
|
|||||||
|
|
||||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(320, 256, 32, 256, 2, 128, 64)
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(320, 256, 32, 256, 2, 128, 64)
|
||||||
|
|
||||||
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 2, 64, 2, 64, 64)
|
||||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 4, 128, 2, 64, 64)
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 4, 128, 2, 64, 64)
|
||||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 8, 256, 2, 64, 64)
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 8, 256, 2, 64, 64)
|
||||||
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 16, 256, 4, 64, 64)
|
GGML_CUDA_FATTN_TILE_CONFIG_CASE(512, 512, 16, 256, 4, 64, 64)
|
||||||
@@ -1308,12 +1312,12 @@ static void launch_fattn_tile_switch_ncols2(ggml_backend_cuda_context & ctx, ggm
|
|||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
if constexpr (DV <= 256) {
|
|
||||||
if (use_gqa_opt && gqa_ratio % 2 == 0) {
|
if (use_gqa_opt && gqa_ratio % 2 == 0) {
|
||||||
launch_fattn_tile_switch_ncols1<DKQ, DV, 2, use_logit_softcap>(ctx, dst);
|
launch_fattn_tile_switch_ncols1<DKQ, DV, 2, use_logit_softcap>(ctx, dst);
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if constexpr (DV <= 256) {
|
||||||
launch_fattn_tile_switch_ncols1<DKQ, DV, 1, use_logit_softcap>(ctx, dst);
|
launch_fattn_tile_switch_ncols1<DKQ, DV, 1, use_logit_softcap>(ctx, dst);
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -99,12 +99,12 @@ static void ggml_cuda_flash_attn_ext_mma_f16_switch_ncols2(ggml_backend_cuda_con
|
|||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
if constexpr (DKQ <= 256) {
|
|
||||||
if (use_gqa_opt && gqa_ratio > 1) {
|
if (use_gqa_opt && gqa_ratio > 1) {
|
||||||
ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1<DKQ, DV, 2>(ctx, dst);
|
ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1<DKQ, DV, 2>(ctx, dst);
|
||||||
return;
|
return;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if constexpr (DKQ <= 256) {
|
||||||
ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1<DKQ, DV, 1>(ctx, dst);
|
ggml_cuda_flash_attn_ext_mma_f16_switch_ncols1<DKQ, DV, 1>(ctx, dst);
|
||||||
} else {
|
} else {
|
||||||
GGML_ABORT("fatal error");
|
GGML_ABORT("fatal error");
|
||||||
|
|||||||
@@ -8,3 +8,4 @@ DECL_FATTN_MMA_F16_CASE(96, 96, 16, 2);
|
|||||||
DECL_FATTN_MMA_F16_CASE(112, 112, 16, 2);
|
DECL_FATTN_MMA_F16_CASE(112, 112, 16, 2);
|
||||||
DECL_FATTN_MMA_F16_CASE(128, 128, 16, 2);
|
DECL_FATTN_MMA_F16_CASE(128, 128, 16, 2);
|
||||||
DECL_FATTN_MMA_F16_CASE(256, 256, 16, 2);
|
DECL_FATTN_MMA_F16_CASE(256, 256, 16, 2);
|
||||||
|
DECL_FATTN_MMA_F16_CASE(512, 512, 16, 2);
|
||||||
|
|||||||
@@ -8,3 +8,4 @@ DECL_FATTN_MMA_F16_CASE(96, 96, 32, 2);
|
|||||||
DECL_FATTN_MMA_F16_CASE(112, 112, 32, 2);
|
DECL_FATTN_MMA_F16_CASE(112, 112, 32, 2);
|
||||||
DECL_FATTN_MMA_F16_CASE(128, 128, 32, 2);
|
DECL_FATTN_MMA_F16_CASE(128, 128, 32, 2);
|
||||||
DECL_FATTN_MMA_F16_CASE(256, 256, 32, 2);
|
DECL_FATTN_MMA_F16_CASE(256, 256, 32, 2);
|
||||||
|
DECL_FATTN_MMA_F16_CASE(512, 512, 32, 2);
|
||||||
|
|||||||
@@ -8,3 +8,4 @@ DECL_FATTN_MMA_F16_CASE(96, 96, 4, 2);
|
|||||||
DECL_FATTN_MMA_F16_CASE(112, 112, 4, 2);
|
DECL_FATTN_MMA_F16_CASE(112, 112, 4, 2);
|
||||||
DECL_FATTN_MMA_F16_CASE(128, 128, 4, 2);
|
DECL_FATTN_MMA_F16_CASE(128, 128, 4, 2);
|
||||||
DECL_FATTN_MMA_F16_CASE(256, 256, 4, 2);
|
DECL_FATTN_MMA_F16_CASE(256, 256, 4, 2);
|
||||||
|
DECL_FATTN_MMA_F16_CASE(512, 512, 4, 2);
|
||||||
|
|||||||
@@ -8,3 +8,4 @@ DECL_FATTN_MMA_F16_CASE(96, 96, 8, 2);
|
|||||||
DECL_FATTN_MMA_F16_CASE(112, 112, 8, 2);
|
DECL_FATTN_MMA_F16_CASE(112, 112, 8, 2);
|
||||||
DECL_FATTN_MMA_F16_CASE(128, 128, 8, 2);
|
DECL_FATTN_MMA_F16_CASE(128, 128, 8, 2);
|
||||||
DECL_FATTN_MMA_F16_CASE(256, 256, 8, 2);
|
DECL_FATTN_MMA_F16_CASE(256, 256, 8, 2);
|
||||||
|
DECL_FATTN_MMA_F16_CASE(512, 512, 8, 2);
|
||||||
|
|||||||
@@ -92,7 +92,7 @@ for ncols in [8, 16, 32, 64]:
|
|||||||
continue
|
continue
|
||||||
if head_size_kq == 320 and ncols2 != 32: # Mistral Small 4
|
if head_size_kq == 320 and ncols2 != 32: # Mistral Small 4
|
||||||
continue
|
continue
|
||||||
if head_size_kq == 512 and ncols2 not in (4, 8): # Gemma 4
|
if head_size_kq == 512 and ncols2 not in (2, 4, 8): # Gemma 4 (+ MTP)
|
||||||
continue
|
continue
|
||||||
if head_size_kq == 576 and ncols2 not in (4, 16, 32): # Deepseek, GLM 4.7 Flash
|
if head_size_kq == 576 and ncols2 not in (4, 16, 32): # Deepseek, GLM 4.7 Flash
|
||||||
continue
|
continue
|
||||||
|
|||||||
Reference in New Issue
Block a user