* test new flash_attn test * rebase and fix to disable subgrou matrices when max_kv_tile == 0 * delete log output * Add i32 support to cpy and enables the all ops test * restore the non target ci tests * comment out of TODO of build-cpu.yml * fix format
This commit is contained in:
parent
a6e2630ff1
commit
0f60f8e691
|
|
@ -2815,11 +2815,25 @@ class ggml_webgpu_shader_lib {
|
||||||
key.common.v_direct &= decisions.use_sg_matrix && key.common.v_type == GGML_TYPE_F16;
|
key.common.v_direct &= decisions.use_sg_matrix && key.common.v_type == GGML_TYPE_F16;
|
||||||
key.use_sg_matrix = decisions.use_sg_matrix;
|
key.use_sg_matrix = decisions.use_sg_matrix;
|
||||||
|
|
||||||
const uint32_t max_kv_tile = ggml_webgpu_flash_attn_max_kv_tile(
|
uint32_t max_kv_tile = ggml_webgpu_flash_attn_max_kv_tile(
|
||||||
context.wg_mem_limit_bytes, decisions.q_tile, decisions.use_sg_matrix ? context.sg_mat_n : 1u,
|
context.wg_mem_limit_bytes, decisions.q_tile, decisions.use_sg_matrix ? context.sg_mat_n : 1u,
|
||||||
key.common.head_dim_qk, key.common.head_dim_v, key.common.has_mask,
|
key.common.head_dim_qk, key.common.head_dim_v, key.common.has_mask,
|
||||||
key.common.k_direct || key.common.v_direct);
|
key.common.k_direct || key.common.v_direct);
|
||||||
GGML_ASSERT(max_kv_tile > 0);
|
|
||||||
|
// WorkGroup storage size isn't enough for some params with subgroup matrices path (ref. https://github.com/ggml-org/llama.cpp/pull/26566)
|
||||||
|
if (max_kv_tile == 0) {
|
||||||
|
GGML_ASSERT(decisions.use_sg_matrix);
|
||||||
|
// switch to flash_attn_reg_tile path
|
||||||
|
decisions.use_sg_matrix = false;
|
||||||
|
decisions.q_tile = GGML_WEBGPU_FLASH_ATTN_TILE_Q_TILE;
|
||||||
|
key.common.k_direct = false;
|
||||||
|
key.common.v_direct = false;
|
||||||
|
key.use_sg_matrix = false;
|
||||||
|
max_kv_tile = ggml_webgpu_flash_attn_max_kv_tile(
|
||||||
|
context.wg_mem_limit_bytes, decisions.q_tile, 1u, key.common.head_dim_qk, key.common.head_dim_v,
|
||||||
|
key.common.has_mask, key.common.k_direct || key.common.v_direct);
|
||||||
|
GGML_ASSERT(max_kv_tile > 0);
|
||||||
|
}
|
||||||
|
|
||||||
decisions.kv_tile = decisions.use_sg_matrix ?
|
decisions.kv_tile = decisions.use_sg_matrix ?
|
||||||
std::min(max_kv_tile, context.sg_mat_n * GGML_WEBGPU_FLASH_ATTN_PREFERRED_KV_SG_TILES) :
|
std::min(max_kv_tile, context.sg_mat_n * GGML_WEBGPU_FLASH_ATTN_PREFERRED_KV_SG_TILES) :
|
||||||
|
|
@ -2993,6 +3007,10 @@ class ggml_webgpu_shader_lib {
|
||||||
defines.push_back("SRC_F16");
|
defines.push_back("SRC_F16");
|
||||||
variant += "_f16";
|
variant += "_f16";
|
||||||
break;
|
break;
|
||||||
|
case GGML_TYPE_I32:
|
||||||
|
defines.push_back("SRC_I32");
|
||||||
|
variant += "_i32";
|
||||||
|
break;
|
||||||
default:
|
default:
|
||||||
GGML_ABORT("Unsupported src type for cpy shader");
|
GGML_ABORT("Unsupported src type for cpy shader");
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -4283,9 +4283,8 @@ static bool ggml_backend_webgpu_device_supports_op(ggml_backend_dev_t dev, const
|
||||||
break;
|
break;
|
||||||
case GGML_OP_CPY:
|
case GGML_OP_CPY:
|
||||||
case GGML_OP_CONT:
|
case GGML_OP_CONT:
|
||||||
supports_op = ((op->type == GGML_TYPE_F32 || op->type == GGML_TYPE_F16) &&
|
supports_op = (op->type == GGML_TYPE_F16 || op->type == GGML_TYPE_F32 || op->type == GGML_TYPE_I32) &&
|
||||||
(src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16)) ||
|
(src0->type == GGML_TYPE_F16 || src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_I32);
|
||||||
(op->type == GGML_TYPE_I32 && src0->type == GGML_TYPE_F32);
|
|
||||||
break;
|
break;
|
||||||
case GGML_OP_SET:
|
case GGML_OP_SET:
|
||||||
supports_op = src0->type == src1->type && src0->type == op->type &&
|
supports_op = src0->type == src1->type && src0->type == op->type &&
|
||||||
|
|
|
||||||
|
|
@ -4,6 +4,8 @@ enable f16;
|
||||||
#define SRC_TYPE f32
|
#define SRC_TYPE f32
|
||||||
#elif defined(SRC_F16)
|
#elif defined(SRC_F16)
|
||||||
#define SRC_TYPE f16
|
#define SRC_TYPE f16
|
||||||
|
#elif defined(SRC_I32)
|
||||||
|
#define SRC_TYPE i32
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
#ifdef DST_F32
|
#ifdef DST_F32
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue