/** FlashAttention */
enum ggml_webgpu_flash_attn_path : uint32_t {
- GGML_WEBGPU_FLASH_ATTN_PATH_SUBGROUP_MATRIX = 0u,
- GGML_WEBGPU_FLASH_ATTN_PATH_TILE = 1u,
- GGML_WEBGPU_FLASH_ATTN_PATH_VEC = 2u,
+ GGML_WEBGPU_FLASH_ATTN_PATH_NONE = 0u,
+ GGML_WEBGPU_FLASH_ATTN_PATH_SUBGROUP_MATRIX = 1u,
+ GGML_WEBGPU_FLASH_ATTN_PATH_TILE = 2u,
+ GGML_WEBGPU_FLASH_ATTN_PATH_VEC = 3u,
};
struct ggml_webgpu_flash_attn_pipeline_key {
};
struct ggml_webgpu_flash_attn_decisions {
- uint32_t path = GGML_WEBGPU_FLASH_ATTN_PATH_SUBGROUP_MATRIX;
+ uint32_t path = GGML_WEBGPU_FLASH_ATTN_PATH_NONE;
uint32_t q_tile = 0;
uint32_t kv_tile = 0;
uint32_t wg_size = 0;
(context.src0->ne[0] % GGML_WEBGPU_FLASH_ATTN_TILE_KV_VEC_WIDTH == 0) &&
(context.src2->ne[0] % GGML_WEBGPU_FLASH_ATTN_TILE_KV_VEC_WIDTH == 0) && !use_vec;
- decisions.path = use_vec ? GGML_WEBGPU_FLASH_ATTN_PATH_VEC :
- use_tile ? GGML_WEBGPU_FLASH_ATTN_PATH_TILE :
- GGML_WEBGPU_FLASH_ATTN_PATH_SUBGROUP_MATRIX;
+ decisions.path = use_vec ? GGML_WEBGPU_FLASH_ATTN_PATH_VEC :
+ use_tile ? GGML_WEBGPU_FLASH_ATTN_PATH_TILE :
+ context.supports_subgroup_matrix ? GGML_WEBGPU_FLASH_ATTN_PATH_SUBGROUP_MATRIX :
+ GGML_WEBGPU_FLASH_ATTN_PATH_NONE;
+
+ if (decisions.path == GGML_WEBGPU_FLASH_ATTN_PATH_NONE) {
+ return decisions;
+ }
const ggml_webgpu_flash_attn_pipeline_key key = ggml_webgpu_flash_attn_make_pipeline_key(context, decisions.path);
decisions.kv_direct = key.kv_direct;
+ const uint32_t max_kv_tile = ggml_webgpu_flash_attn_max_kv_tile(context, key);
+ // invalidate if even the smallest kv_tile doesn't fit in shared memory
+ if (max_kv_tile == 0) {
+ decisions.path = GGML_WEBGPU_FLASH_ATTN_PATH_NONE;
+ return decisions;
+ }
if (decisions.path == GGML_WEBGPU_FLASH_ATTN_PATH_VEC) {
- const uint32_t min_kv_tile = ggml_webgpu_flash_attn_max_kv_tile(context, key);
- decisions.q_tile = 1u;
- decisions.kv_tile = std::max(8u, std::min(32u, min_kv_tile));
- decisions.kv_tile = (decisions.kv_tile / 8u) * 8u;
- decisions.wg_size = std::max(1u, std::min<uint32_t>(32u, context.max_subgroup_size));
+ decisions.q_tile = 1u;
+ decisions.kv_tile = std::max(8u, std::min(32u, max_kv_tile));
+ decisions.kv_tile = (decisions.kv_tile / 8u) * 8u;
+ decisions.wg_size = std::max(1u, std::min<uint32_t>(32u, context.max_subgroup_size));
if (decisions.kv_direct) {
decisions.kv_tile = std::min(decisions.kv_tile, GGML_WEBGPU_KV_SEQ_PAD);
while (GGML_WEBGPU_KV_SEQ_PAD % decisions.kv_tile != 0) {
decisions.q_tile =
decisions.path == GGML_WEBGPU_FLASH_ATTN_PATH_TILE ? GGML_WEBGPU_FLASH_ATTN_TILE_Q_TILE : context.sg_mat_m;
decisions.kv_tile = decisions.path == GGML_WEBGPU_FLASH_ATTN_PATH_TILE ?
- std::min(64u, ggml_webgpu_flash_attn_max_kv_tile(context, key)) :
- std::min(ggml_webgpu_flash_attn_max_kv_tile(context, key),
- context.sg_mat_n * GGML_WEBGPU_FLASH_ATTN_PREFERRED_KV_SG_TILES);
+ std::min(64u, max_kv_tile) :
+ std::min(max_kv_tile, context.sg_mat_n * GGML_WEBGPU_FLASH_ATTN_PREFERRED_KV_SG_TILES);
decisions.wg_size = decisions.path == GGML_WEBGPU_FLASH_ATTN_PATH_TILE ?
GGML_WEBGPU_FLASH_ATTN_PREFERRED_WG_SIZE :
std::max(context.max_subgroup_size, GGML_WEBGPU_FLASH_ATTN_PREFERRED_WG_SIZE);
context.sg_mat_n;
}
}
-
return decisions;
}
if (key.src_type == GGML_TYPE_Q1_0) {
defines.push_back("BLOCK_SIZE=128u");
} else if ((key.src_type >= GGML_TYPE_Q4_0 && key.src_type <= GGML_TYPE_Q8_1) ||
- key.src_type == GGML_TYPE_IQ4_NL) {
+ key.src_type == GGML_TYPE_IQ4_NL) {
defines.push_back("BLOCK_SIZE=32u");
} else if (key.src_type >= GGML_TYPE_Q2_K) {
defines.push_back("BLOCK_SIZE=256u");
size_t storage_offset_alignment) {
const ggml_webgpu_flash_attn_decisions decisions =
ggml_webgpu_flash_attn_get_decisions(context, storage_offset_alignment);
+ GGML_ASSERT(decisions.path != GGML_WEBGPU_FLASH_ATTN_PATH_NONE);
ggml_webgpu_flash_attn_pipeline_key key = ggml_webgpu_flash_attn_make_pipeline_key(context, decisions.path);
auto it = flash_attn_pipelines.find(key);
if (it != flash_attn_pipelines.end()) {