From f09a97cf6dddc821bc029c6043cd061aff8be66d Mon Sep 17 00:00:00 2001 From: Masashi Yoshimura Date: Mon, 10 Aug 2026 15:29:41 +0900 Subject: [PATCH] ggml-webgpu : refactor several wgsl files and simplify flash_attn wgsl. (llama/26134) --- .../ggml-webgpu/ggml-webgpu-shader-lib.hpp | 31 ++- ggml/src/ggml-webgpu/ggml-webgpu.cpp | 36 +-- ggml/src/ggml-webgpu/wgsl-shaders/concat.wgsl | 1 - ggml/src/ggml-webgpu/wgsl-shaders/conv2d.wgsl | 50 +--- .../ggml-webgpu/wgsl-shaders/conv2d_dw.wgsl | 53 +--- .../ggml-webgpu/wgsl-shaders/flash_attn.wgsl | 161 +---------- .../wgsl-shaders/flash_attn_decls.tmpl | 134 +++++++++ .../flash_attn_quant_staging.tmpl | 83 ------ .../wgsl-shaders/flash_attn_staging.tmpl | 136 +++++++++ .../wgsl-shaders/flash_attn_tile.wgsl | 180 +----------- .../wgsl-shaders/flash_attn_vec_split.wgsl | 260 ++---------------- ggml/src/ggml-webgpu/wgsl-shaders/im2col.wgsl | 36 +-- .../wgsl-shaders/rms_norm_mul.wgsl | 1 - .../ggml-webgpu/wgsl-shaders/row_norm.wgsl | 1 - .../ggml-webgpu/wgsl-shaders/soft_max.wgsl | 69 ++--- .../ggml-webgpu/wgsl-shaders/solve_tri.wgsl | 1 - .../ggml-webgpu/wgsl-shaders/ssm_scan.wgsl | 1 - 17 files changed, 372 insertions(+), 862 deletions(-) create mode 100644 ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_decls.tmpl delete mode 100644 ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_quant_staging.tmpl create mode 100644 ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_staging.tmpl diff --git a/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp b/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp index 66c1c3c89..35a55ecaf 100644 --- a/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp +++ b/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp @@ -3221,17 +3221,17 @@ class ggml_webgpu_shader_lib { auto push_type_defines = [&](const char * prefix, ggml_type type) { std::string s_prefix = prefix; if (type == GGML_TYPE_F32) { - defines.push_back(s_prefix + "_F32"); + defines.push_back(s_prefix + "=f32"); } else if (type == GGML_TYPE_F16) { - defines.push_back(s_prefix + "_F16"); + defines.push_back(s_prefix + "=f16"); } else { GGML_ABORT("Unsupported type for CONV_2D shader"); } }; - push_type_defines("WEIGHT", key.weight_type); - push_type_defines("INPUT", key.input_type); - push_type_defines("OUTPUT", key.output_type); + push_type_defines("WEIGHT_TYPE", key.weight_type); + push_type_defines("INPUT_TYPE", key.input_type); + push_type_defines("OUTPUT_TYPE", key.output_type); defines.push_back(std::string("WG_SIZE=") + std::to_string(context.max_wg_size)); @@ -3263,17 +3263,18 @@ class ggml_webgpu_shader_lib { auto push_type_defines = [&](const char * prefix, ggml_type type) { std::string s_prefix = prefix; if (type == GGML_TYPE_F32) { - defines.push_back(s_prefix + "_F32"); + defines.push_back(s_prefix + "=f32"); } else if (type == GGML_TYPE_F16) { - defines.push_back(s_prefix + "_F16"); + defines.push_back(s_prefix + "=f16"); } else { - GGML_ABORT("Unsupported type for CONV_2D_DW shader"); + GGML_ABORT("Unsupported type for CONV_2D shader"); } }; - push_type_defines("WEIGHT", key.weight_type); - push_type_defines("INPUT", key.input_type); - push_type_defines("OUTPUT", key.output_type); + push_type_defines("WEIGHT_TYPE", key.weight_type); + push_type_defines("INPUT_TYPE", key.input_type); + push_type_defines("OUTPUT_TYPE", key.output_type); + if (whcn) { defines.push_back("WHCN"); } @@ -3304,16 +3305,16 @@ class ggml_webgpu_shader_lib { auto push_type_defines = [&](const char * prefix, ggml_type type) { std::string s_prefix = prefix; if (type == GGML_TYPE_F32) { - defines.push_back(s_prefix + "_F32"); + defines.push_back(s_prefix + "=f32"); } else if (type == GGML_TYPE_F16) { - defines.push_back(s_prefix + "_F16"); + defines.push_back(s_prefix + "=f16"); } else { GGML_ABORT("Unsupported type for IM2COL shader"); } }; - push_type_defines("INPUT", key.input_type); - push_type_defines("OUTPUT", key.output_type); + push_type_defines("INPUT_TYPE", key.input_type); + push_type_defines("OUTPUT_TYPE", key.output_type); defines.push_back(std::string("WG_SIZE=") + std::to_string(context.max_wg_size)); diff --git a/ggml/src/ggml-webgpu/ggml-webgpu.cpp b/ggml/src/ggml-webgpu/ggml-webgpu.cpp index c001cda7d..ba4b91695 100644 --- a/ggml/src/ggml-webgpu/ggml-webgpu.cpp +++ b/ggml/src/ggml-webgpu/ggml-webgpu.cpp @@ -930,7 +930,6 @@ static webgpu_encoded_op ggml_webgpu_solve_tri(webgpu_context & ctx, (uint32_t) src1->ne[0], (uint32_t) dst->ne[2], - (uint32_t) dst->ne[3], }; std::vector entries = { @@ -1039,7 +1038,6 @@ static webgpu_encoded_op ggml_webgpu_conv_2d_dw(webgpu_context & ctx, (uint32_t) ggml_nelements(dst), (uint32_t) dst->ne[2], - (uint32_t) dst->ne[3], (uint32_t) dst->ne[0], (uint32_t) dst->ne[1], (uint32_t) src1->ne[0], @@ -1328,7 +1326,6 @@ static webgpu_encoded_op ggml_webgpu_ssm_scan(webgpu_context & ctx, (uint32_t) src0->ne[2], (uint32_t) src4->ne[1], (uint32_t) src1->ne[2], - (uint32_t) src1->ne[3], (uint32_t) ggml_nelements(src1), }; @@ -1921,25 +1918,20 @@ static bool ggml_webgpu_flash_attn_use_vec_path(const webgpu_global_context & gl const ggml_tensor * K, const ggml_tensor * V) { const size_t storage_offset_alignment = global_ctx->capabilities.limits.minStorageBufferOffsetAlignment; - const bool k_float_vec4_aligned = (K->type != GGML_TYPE_F16 && K->type != GGML_TYPE_F32) || - ggml_webgpu_flash_attn_float_vec4_aligned(K, storage_offset_alignment); - const bool v_float_vec4_aligned = (V->type != GGML_TYPE_F16 && V->type != GGML_TYPE_F32) || - ggml_webgpu_flash_attn_float_vec4_aligned(V, storage_offset_alignment); - const bool k_vec_type_supported = - K->type == GGML_TYPE_F32 || K->type == GGML_TYPE_F16 || K->type == GGML_TYPE_Q4_0 || K->type == GGML_TYPE_Q8_0; - const bool v_vec_type_supported = - V->type == GGML_TYPE_F32 || V->type == GGML_TYPE_F16 || V->type == GGML_TYPE_Q4_0 || V->type == GGML_TYPE_Q8_0; - const uint32_t k_vec_head_align = (K->type == GGML_TYPE_F32 || K->type == GGML_TYPE_F16) ? - GGML_WEBGPU_FLASH_ATTN_TILE_KV_VEC_WIDTH : - (uint32_t) ggml_blck_size(K->type); - const uint32_t v_vec_head_align = (V->type == GGML_TYPE_F32 || V->type == GGML_TYPE_F16) ? - GGML_WEBGPU_FLASH_ATTN_TILE_KV_VEC_WIDTH : - (uint32_t) ggml_blck_size(V->type); - const bool kv_vec_head_dims_aligned = Q->ne[0] % k_vec_head_align == 0 && V->ne[0] % v_vec_head_align == 0; + + const bool k_float_vec4_aligned = (K->type != GGML_TYPE_F16 && K->type != GGML_TYPE_F32) || + ggml_webgpu_flash_attn_float_vec4_aligned(K, storage_offset_alignment); + const bool v_float_vec4_aligned = (V->type != GGML_TYPE_F16 && V->type != GGML_TYPE_F32) || + ggml_webgpu_flash_attn_float_vec4_aligned(V, storage_offset_alignment); + + const uint32_t k_vec_head_align = + ggml_is_quantized(K->type) ? ggml_blck_size(K->type) : GGML_WEBGPU_FLASH_ATTN_TILE_KV_VEC_WIDTH; + const uint32_t v_vec_head_align = + ggml_is_quantized(V->type) ? ggml_blck_size(V->type) : GGML_WEBGPU_FLASH_ATTN_TILE_KV_VEC_WIDTH; + const bool kv_vec_head_dims_aligned = Q->ne[0] % k_vec_head_align == 0 && V->ne[0] % v_vec_head_align == 0; return global_ctx->capabilities.supports_subgroups && (Q->ne[1] < GGML_WEBGPU_FLASH_ATTN_VEC_MAX_SEQ_LEN) && - kv_vec_head_dims_aligned && k_vec_type_supported && v_vec_type_supported && k_float_vec4_aligned && - v_float_vec4_aligned; + kv_vec_head_dims_aligned && k_float_vec4_aligned && v_float_vec4_aligned; } static ggml_webgpu_flash_attn_op ggml_webgpu_flash_attn_prepare(webgpu_context & ctx, @@ -2514,7 +2506,6 @@ static webgpu_encoded_op ggml_webgpu_concat(webgpu_context & ctx, (uint32_t) dst->ne[0], (uint32_t) dst->ne[1], (uint32_t) dst->ne[2], - (uint32_t) dst->ne[3], dim, (uint32_t) src0->ne[dim] }; @@ -2610,7 +2601,6 @@ static std::optional ggml_webgpu_rms_norm_mul(webgpu_context (uint32_t) dst->ne[0], (uint32_t) dst->ne[1], (uint32_t) dst->ne[2], - (uint32_t) dst->ne[3], ggml_webgpu_u32_from_f32(ggml_get_op_params_f32(rn_dst, 0)) // epsilon, treated as f32 in the shader }; @@ -2666,7 +2656,6 @@ static webgpu_encoded_op ggml_webgpu_row_norm(webgpu_context & ctx, ggml_tensor (uint32_t) src->ne[0], (uint32_t) src->ne[1], (uint32_t) src->ne[2], - (uint32_t) src->ne[3], ggml_webgpu_u32_from_f32(ggml_get_op_params_f32(dst, 0)) // epsilon, treated as f32 in the shader }; @@ -2925,7 +2914,6 @@ static webgpu_encoded_op ggml_webgpu_soft_max(webgpu_context & ctx, (uint32_t) (dst->nb[1] / ggml_type_size(dst->type)), (uint32_t) (dst->nb[2] / ggml_type_size(dst->type)), (uint32_t) (dst->nb[3] / ggml_type_size(dst->type)), - (uint32_t) ggml_nelements(dst), (uint32_t) src0->ne[0], (uint32_t) src0->ne[1], (uint32_t) src0->ne[2], diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/concat.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/concat.wgsl index eb901bf05..7ccad73f4 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/concat.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/concat.wgsl @@ -18,7 +18,6 @@ struct Params { ne0: u32, ne1: u32, ne2: u32, - ne3: u32, dim: u32, src0_nedim: u32 diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/conv2d.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/conv2d.wgsl index 9eb131dc2..38c714ba5 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/conv2d.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/conv2d.wgsl @@ -2,25 +2,11 @@ enable f16; @group(0) @binding(0) -#if defined(WEIGHT_F32) -var weights: array; -#elif defined(WEIGHT_F16) -var weights: array; -#endif - +var weights: array; @group(0) @binding(1) -#if defined(INPUT_F32) -var input: array; -#elif defined(INPUT_F16) -var input: array; -#endif - +var input: array; @group(0) @binding(2) -#if defined(OUTPUT_F32) -var output: array; -#elif defined(OUTPUT_F16) -var output: array; -#endif +var output: array; struct Params { offset_w: u32, @@ -50,30 +36,6 @@ struct Params { @group(0) @binding(3) var params: Params; -fn load_weight(idx: u32) -> f32 { - #if defined(WEIGHT_F32) - return weights[idx]; - #elif defined(WEIGHT_F16) - return f32(weights[idx]); - #endif -} - -fn load_input(idx: u32) -> f32 { - #if defined(INPUT_F32) - return input[idx]; - #elif defined(INPUT_F16) - return f32(input[idx]); - #endif -} - -fn store_output(idx: u32, val: f32) { - #if defined(OUTPUT_F32) - output[idx] = val; - #elif defined(OUTPUT_F16) - output[idx] = f16(val); - #endif -} - fn ceil_div_u32(x: u32, y: u32) -> u32 { return (x + y - 1) / y; } @@ -136,7 +98,7 @@ fn main( // entire receptive field is out of bounds if (kw_begin >= kw_end || kh_begin >= kh_end) { let out_idx = params.offset_o + ow * params.so0 + oh * params.so1 + oc * params.so2 + n * params.so3; - store_output(out_idx, 0.0); + output[out_idx] = OUTPUT_TYPE(0.0); return; } @@ -155,11 +117,11 @@ fn main( let iw = u32(ow_base + i32(kw * params.d0)); let w_idx = w_row_base + kw * params.sw0; let in_idx = in_row_base + iw * params.si0; - sum += load_weight(w_idx) * load_input(in_idx); + sum += f32(weights[w_idx]) * f32(input[in_idx]); } } } let out_idx = params.offset_o + ow * params.so0 + oh * params.so1 + oc * params.so2 + n * params.so3; - store_output(out_idx, sum); + output[out_idx] = OUTPUT_TYPE(sum); } diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/conv2d_dw.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/conv2d_dw.wgsl index 42d6f027c..fc028e429 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/conv2d_dw.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/conv2d_dw.wgsl @@ -6,25 +6,11 @@ enable f16; // weight (src0) is [KW,KH,1,C]; output matches the input layout. @group(0) @binding(0) -#if defined(WEIGHT_F32) -var weights: array; -#elif defined(WEIGHT_F16) -var weights: array; -#endif - +var weights: array; @group(0) @binding(1) -#if defined(INPUT_F32) -var input: array; -#elif defined(INPUT_F16) -var input: array; -#endif - +var input: array; @group(0) @binding(2) -#if defined(OUTPUT_F32) -var output: array; -#elif defined(OUTPUT_F16) -var output: array; -#endif +var output: array; struct Params { offset_w: u32, @@ -33,7 +19,6 @@ struct Params { ne: u32, channels: u32, - batches: u32, dst_w: u32, dst_h: u32, src_w: u32, src_h: u32, knl_w: u32, knl_h: u32, @@ -46,28 +31,6 @@ struct Params { @group(0) @binding(3) var params: Params; -fn load_weight(idx: u32) -> f32 { - #if defined(WEIGHT_F32) - return weights[idx]; - #elif defined(WEIGHT_F16) - return f32(weights[idx]); - #endif -} -fn load_input(idx: u32) -> f32 { - #if defined(INPUT_F32) - return input[idx]; - #elif defined(INPUT_F16) - return f32(input[idx]); - #endif -} -fn store_output(idx: u32, val: f32) { - #if defined(OUTPUT_F32) - output[idx] = val; - #elif defined(OUTPUT_F16) - output[idx] = f16(val); - #endif -} - #if defined(WHCN) // Input/output/kernel contiguous in [W, H, C, N] order (kernel [KW,KH,C]). fn conv_2d_dw(idx: u32) -> f32 { @@ -89,8 +52,8 @@ fn conv_2d_dw(idx: u32) -> f32 { for (var kx: u32 = 0u; kx < params.knl_w; kx += 1u) { let src_x = i32(dst_x) * params.stride_x + i32(kx) * params.dilation_x - params.pad_x; if (src_x < 0 || src_x >= i32(params.src_w)) { continue; } - let v = load_input(src_i + u32(src_y) * params.src_w + u32(src_x)); - let k = load_weight(knl_i + ky * params.knl_w + kx); + let v = f32(input[src_i + u32(src_y) * params.src_w + u32(src_x)]); + let k = f32(weights[knl_i + ky * params.knl_w + kx]); sum += v * k; } } @@ -117,8 +80,8 @@ fn conv_2d_dw(idx: u32) -> f32 { for (var kx: u32 = 0u; kx < params.knl_w; kx += 1u) { let src_x = i32(dst_x) * params.stride_x + i32(kx) * params.dilation_x - params.pad_x; if (src_x < 0 || src_x >= i32(params.src_w)) { continue; } - let v = load_input(src_i + u32(src_y) * src_row + u32(src_x) * params.channels + c); - let k = load_weight(params.offset_w + ky * knl_row + kx * params.channels + c); + let v = f32(input[src_i + u32(src_y) * src_row + u32(src_x) * params.channels + c]); + let k = f32(weights[params.offset_w + ky * knl_row + kx * params.channels + c]); sum += v * k; } } @@ -133,5 +96,5 @@ fn main( ) { let idx = gid.x + (num_wg.x * u32(WG_SIZE)) * gid.y; if (idx >= params.ne) { return; } - store_output(params.offset_o + idx, conv_2d_dw(idx)); + output[params.offset_o + idx] = OUTPUT_TYPE(conv_2d_dw(idx)); } diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn.wgsl index 75f33e68a..d5bf2af8d 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn.wgsl @@ -7,32 +7,18 @@ enable chromium_experimental_subgroup_matrix; #define BYTE_HELPERS #include "common_decls.tmpl" -#ifdef K_F32 -#define K_TYPE f32 -#elif defined(K_Q4_0) || defined(K_Q8_0) -#define K_TYPE u32 -#else -#define K_TYPE f16 -#endif - -#ifdef V_F32 -#define V_TYPE f32 -#elif defined(V_Q4_0) || defined(V_Q8_0) -#define V_TYPE u32 -#else -#define V_TYPE f16 -#endif +#define FLASH_ATTN_SCALAR_KV +#include "flash_attn_decls.tmpl" // Default values +// The actual values are defined in shader-lib. #define HEAD_DIM_QK 64 #define HEAD_DIM_V 64 - // The number of rows/columns/k in a subgroup matrix. MxK * KxN = MxN // Note that the "K" here does not correspond to the K in attention's Q/K/V, it's just the common dimension. #define SG_MAT_M 8 #define SG_MAT_N 8 #define SG_MAT_K 8 - // Each workgroup processes one subgroup matrix of Q rows #define Q_TILE SG_MAT_M #define KV_TILE 16 @@ -41,104 +27,13 @@ enable chromium_experimental_subgroup_matrix; // Number of subgroup-matrix-width blocks that span the KV tile. SG_MAT_N must divide KV_TILE. #define KV_BLOCKS (KV_TILE / SG_MAT_N) -struct Params { - offset_q: u32, - offset_k: u32, - offset_v: u32, - offset_mask: u32, - offset_sinks: u32, - offset_dst: u32, - - // shapes of Q/K/V - n_heads: u32, - seq_len_q: u32, - seq_len_kv: u32, - - // strides (in elements) - stride_q1: u32, - stride_q2: u32, - stride_q3: u32, - stride_k1: u32, - stride_k2: u32, - stride_k3: u32, - stride_v1: u32, - stride_v2: u32, - stride_v3: u32, - stride_mask3: u32, - - // repeat factors for K/V, e.g., MHA vs. MQA vs. GQA - q_per_kv: u32, - - // softmax params - scale: f32, - max_bias: f32, - logit_softcap: f32, - n_head_log2: f32, - m0: f32, - m1: f32, -}; - -@group(0) @binding(0) var Q: array; -#ifdef KV_OVERLAP -@group(0) @binding(1) var K: array; -#define V K -#else -@group(0) @binding(1) var K: array; -@group(0) @binding(2) var V: array; -#endif - -#if defined(MASK) && defined(SINKS) -#ifdef KV_OVERLAP -@group(0) @binding(2) var mask: array; -@group(0) @binding(3) var sinks: array; -#define DST_BINDING 4 -#define PARAMS_BINDING 5 -#else -@group(0) @binding(3) var mask: array; -@group(0) @binding(4) var sinks: array; -#define DST_BINDING 5 -#define PARAMS_BINDING 6 -#endif -#elif defined(MASK) -#ifdef KV_OVERLAP -@group(0) @binding(2) var mask: array; -#define DST_BINDING 3 -#define PARAMS_BINDING 4 -#else -@group(0) @binding(3) var mask: array; -#define DST_BINDING 4 -#define PARAMS_BINDING 5 -#endif -#elif defined(SINKS) -#ifdef KV_OVERLAP -@group(0) @binding(2) var sinks: array; -#define DST_BINDING 3 -#define PARAMS_BINDING 4 -#else -@group(0) @binding(3) var sinks: array; -#define DST_BINDING 4 -#define PARAMS_BINDING 5 -#endif -#else -#ifdef KV_OVERLAP -#define DST_BINDING 2 -#define PARAMS_BINDING 3 -#else -#define DST_BINDING 3 -#define PARAMS_BINDING 4 -#endif -#endif - -@group(0) @binding(DST_BINDING) var dst: array>; -@group(0) @binding(PARAMS_BINDING) var params: Params; - -// Just a very small float value. -const FLOAT_MIN: f32 = -1.0e9; - // The number of Q rows processed per workgroup var q_shmem: array; #if !defined(K_DIRECT) || !defined(V_DIRECT) +#define STAGING_SHMEM kv_shmem +#define STAGING_OUT_TYPE f16 +#include "flash_attn_staging.tmpl" const kv_shmem_size = KV_TILE * max(HEAD_DIM_QK, HEAD_DIM_V); // we can reuse the same shmem for K and V since we only need one at a time var kv_shmem: array; @@ -175,50 +70,6 @@ fn calc_softmax_term(kv_idx: u32, q_tile_row: u32, slope: f32) -> f32 { return v; } -fn load_f32x4(buf: ptr>, read_write>, scalar_index: u32) -> vec4 { - return (*buf)[scalar_index >> 2u]; -} - -fn load_kx4(buf: ptr>, read_write>, scalar_index: u32) -> vec4 { - return (*buf)[scalar_index >> 2u]; -} - -#if !defined(K_DIRECT) || !defined(V_DIRECT) -#define QUANT_SHMEM kv_shmem -#define QUANT_OUT_TYPE f16 -#include "flash_attn_quant_staging.tmpl" - -#if !defined(K_DIRECT) && !defined(K_Q4_0) && !defined(K_Q8_0) -fn load_k_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, k_head_offset: u32) { - for (var elem_idx = local_x; elem_idx < KV_TILE * HEAD_DIM_QK; elem_idx += WG_SIZE) { - let k_row = elem_idx / HEAD_DIM_QK; - let k_col = elem_idx % HEAD_DIM_QK; - let global_k_row = kv_tile + k_row; - let global_k_row_offset = k_head_offset + global_k_row * params.stride_k1; - kv_shmem[elem_idx] = f16(select( - 0.0, - K[global_k_row_offset + k_col], - global_k_row < params.seq_len_kv && k_col < HEAD_DIM_QK)); - } -} -#endif - -#if !defined(V_DIRECT) && !defined(V_Q4_0) && !defined(V_Q8_0) -fn load_v_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, v_head_offset: u32) { - for (var elem_idx = local_x; elem_idx < KV_TILE * HEAD_DIM_V; elem_idx += WG_SIZE) { - let v_row = elem_idx / HEAD_DIM_V; - let v_col = elem_idx % HEAD_DIM_V; - let global_v_row = kv_tile + v_row; - let global_v_row_offset = v_head_offset + global_v_row * params.stride_v1; - kv_shmem[elem_idx] = f16(select( - 0.0, - V[global_v_row_offset + v_col], - global_v_row < params.seq_len_kv && v_col < HEAD_DIM_V)); - } -} -#endif -#endif - @compute @workgroup_size(WG_SIZE) fn main(@builtin(workgroup_id) wg_id: vec3, @builtin(local_invocation_id) local_id: vec3, diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_decls.tmpl b/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_decls.tmpl new file mode 100644 index 000000000..48a79b6ce --- /dev/null +++ b/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_decls.tmpl @@ -0,0 +1,134 @@ +#ifdef Q_F32 +#define Q_TYPE f32 +#else +#define Q_TYPE f16 +#endif + +#ifdef K_F32 +#define K_TYPE f32 +#elif defined(K_Q4_0) || defined(K_Q8_0) +#define K_TYPE u32 +#else +#define K_TYPE f16 +#endif + +#ifdef V_F32 +#define V_TYPE f32 +#elif defined(V_Q4_0) || defined(V_Q8_0) +#define V_TYPE u32 +#else +#define V_TYPE f16 +#endif + +#ifdef DST_F32 +#define DST_TYPE f32 +#else +#define DST_TYPE f16 +#endif + +#if defined(FLASH_ATTN_SCALAR_KV) || defined(K_Q4_0) || defined(K_Q8_0) +#define K_STORAGE_TYPE K_TYPE +#else +#define K_STORAGE_TYPE vec4 +#endif + +#if defined(FLASH_ATTN_SCALAR_KV) || defined(V_Q4_0) || defined(V_Q8_0) +#define V_STORAGE_TYPE V_TYPE +#else +#define V_STORAGE_TYPE vec4 +#endif + +// Just a very small float value. +const FLOAT_MIN: f32 = -1.0e9; + +struct Params { + offset_q: u32, + offset_k: u32, + offset_v: u32, + offset_mask: u32, + offset_sinks: u32, + offset_dst: u32, + + // shapes of Q/K/V + n_heads: u32, + seq_len_q: u32, + seq_len_kv: u32, + + // strides (in elements) + stride_q1: u32, + stride_q2: u32, + stride_q3: u32, + stride_k1: u32, + stride_k2: u32, + stride_k3: u32, + stride_v1: u32, + stride_v2: u32, + stride_v3: u32, + stride_mask3: u32, + + // repeat factors for K/V, e.g., MHA vs. MQA vs. GQA + q_per_kv: u32, + + // softmax params + scale: f32, + max_bias: f32, + logit_softcap: f32, + n_head_log2: f32, + m0: f32, + m1: f32, + +#ifdef FLASH_ATTN_VEC_SPLIT +#ifdef BLK + blk_base: u32, + blk_nblk0: u32, + blk_nblk1: u32, +#endif + + tmp_data_base: u32, + tmp_stats_base: u32, + nwg: u32, +#endif +}; + +@group(0) @binding(0) var Q: array; +@group(0) @binding(1) var K: array; +#ifdef KV_OVERLAP +#define V K +#define MASK_BINDING 2 +#else +@group(0) @binding(2) var V: array; +#define MASK_BINDING 3 +#endif // KV_OVERLAP + +#ifdef MASK +@group(0) @binding(MASK_BINDING) var mask: array; +#define SINKS_BINDING (MASK_BINDING + 1) +#else +#define SINKS_BINDING MASK_BINDING +#endif + +#ifdef SINKS +@group(0) @binding(SINKS_BINDING) var sinks: array; +#define BLK_BINDING (SINKS_BINDING + 1) +#else +#define BLK_BINDING SINKS_BINDING +#endif + +#ifdef FLASH_ATTN_VEC_SPLIT +#ifdef BLK +@group(0) @binding(BLK_BINDING) var blk: array; +#define TMP_BINDING (BLK_BINDING + 1) +#else +#define TMP_BINDING BLK_BINDING +#endif + +@group(0) @binding(TMP_BINDING) var tmp: array; +#define DST_BINDING (TMP_BINDING + 1) +#else +#define DST_BINDING BLK_BINDING +#endif // FLASH_ATTN_VEC_SPLIT + +@group(0) @binding(DST_BINDING) var dst: array>; + +#define PARAMS_BINDING (DST_BINDING + 1) +@group(0) @binding(PARAMS_BINDING) var params: Params; diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_quant_staging.tmpl b/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_quant_staging.tmpl deleted file mode 100644 index 1c23260df..000000000 --- a/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_quant_staging.tmpl +++ /dev/null @@ -1,83 +0,0 @@ -#include "quant_inner_loops.tmpl" - -#define BLOCK_SIZE 32 -#define BLOCKS_K ((HEAD_DIM_QK + BLOCK_SIZE - 1) / BLOCK_SIZE) -#define BLOCKS_V ((HEAD_DIM_V + BLOCK_SIZE - 1) / BLOCK_SIZE) - -#if defined(K_Q4_0) -#define K_NQ 16 -#define K_BLOCK_SIZE_BYTES 18u -#define K_BYTES_PER_THREAD 8u -#define K_BYTES_PER_INNER_LOOP 4u -#elif defined(K_Q8_0) -#define K_NQ 16 -#define K_BLOCK_SIZE_BYTES 34u -#define K_BYTES_PER_THREAD 16u -#define K_BYTES_PER_INNER_LOOP 4u -#endif - -#if defined(V_Q4_0) -#define V_NQ 16 -#define V_BLOCK_SIZE_BYTES 18u -#define V_BYTES_PER_THREAD 8u -#define V_BYTES_PER_INNER_LOOP 4u -#elif defined(V_Q8_0) -#define V_NQ 16 -#define V_BLOCK_SIZE_BYTES 34u -#define V_BYTES_PER_THREAD 16u -#define V_BYTES_PER_INNER_LOOP 4u -#endif - -#if defined(K_Q4_0) || defined(K_Q8_0) -fn load_k_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, k_head_offset: u32) { - for (var elem_idx = local_x * K_NQ; elem_idx < kv_count * HEAD_DIM_QK; elem_idx += WG_SIZE * K_NQ) { - let blck_idx = elem_idx / BLOCK_SIZE; - let block_offset = (elem_idx % BLOCK_SIZE) / K_NQ; - let k_row = blck_idx / BLOCKS_K; - let global_k_row = kv_tile + k_row; - let block_k = blck_idx % BLOCKS_K; - let row_offset = k_row * HEAD_DIM_QK; - let global_block_idx = k_head_offset + global_k_row * params.stride_k1 + block_k; - let block_byte_base = global_block_idx * K_BLOCK_SIZE_BYTES; - let d = f16_from_u16(load_k_u16_at(block_byte_base)); - let thread_byte_offset = block_offset * K_BYTES_PER_THREAD; - let shmem_idx = row_offset + block_k * BLOCK_SIZE + thread_byte_offset; - for (var j = 0u; j < K_BYTES_PER_THREAD / K_BYTES_PER_INNER_LOOP; j += 1u) { - let q_byte_offset = block_byte_base + 2u + thread_byte_offset + j * K_BYTES_PER_INNER_LOOP; - let q_packed = load_k_u32_at(q_byte_offset); -#if defined(K_Q4_0) - dequant_q4_0_packed_to_shmem(q_packed, d, shmem_idx + j * K_BYTES_PER_INNER_LOOP); -#elif defined(K_Q8_0) - dequant_q8_0_packed_to_shmem(q_packed, d, shmem_idx + j * K_BYTES_PER_INNER_LOOP); -#endif - } - } -} -#endif - -#if defined(V_Q4_0) || defined(V_Q8_0) -fn load_v_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, v_head_offset: u32) { - for (var elem_idx = local_x * V_NQ; elem_idx < kv_count * HEAD_DIM_V; elem_idx += WG_SIZE * V_NQ) { - let blck_idx = elem_idx / BLOCK_SIZE; - let block_offset = (elem_idx % BLOCK_SIZE) / V_NQ; - let v_row = blck_idx / BLOCKS_V; - let global_v_row = kv_tile + v_row; - let block_k = blck_idx % BLOCKS_V; - let row_offset = v_row * HEAD_DIM_V; - let global_block_idx = v_head_offset + global_v_row * params.stride_v1 + block_k; - let block_byte_base = global_block_idx * V_BLOCK_SIZE_BYTES; - let d = f16_from_u16(load_v_u16_at(block_byte_base)); - let thread_byte_offset = block_offset * V_BYTES_PER_THREAD; - let shmem_idx = row_offset + block_k * BLOCK_SIZE + thread_byte_offset; - for (var j = 0u; j < V_BYTES_PER_THREAD / V_BYTES_PER_INNER_LOOP; j += 1u) { - let q_byte_offset = block_byte_base + 2u + thread_byte_offset + j * V_BYTES_PER_INNER_LOOP; - let q_packed = load_v_u32_at(q_byte_offset); -#if defined(V_Q4_0) - dequant_q4_0_packed_to_shmem(q_packed, d, shmem_idx + j * V_BYTES_PER_INNER_LOOP); -#elif defined(V_Q8_0) - dequant_q8_0_packed_to_shmem(q_packed, d, shmem_idx + j * V_BYTES_PER_INNER_LOOP); -#endif - } - } -} -#endif diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_staging.tmpl b/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_staging.tmpl new file mode 100644 index 000000000..457df07ff --- /dev/null +++ b/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_staging.tmpl @@ -0,0 +1,136 @@ +#if defined(K_Q4_0) || defined(K_Q8_0) || defined(V_Q4_0) || defined(V_Q8_0) +#define QUANT_SHMEM STAGING_SHMEM +#define QUANT_OUT_TYPE STAGING_OUT_TYPE +#include "quant_inner_loops.tmpl" +#undef QUANT_SHMEM +#undef QUANT_OUT_TYPE +#define BLOCK_SIZE 32 +#define BLOCKS_K ((HEAD_DIM_QK + BLOCK_SIZE - 1) / BLOCK_SIZE) +#define BLOCKS_V ((HEAD_DIM_V + BLOCK_SIZE - 1) / BLOCK_SIZE) +#endif + +#if defined(K_Q4_0) +#define K_NQ 16 +#define K_BLOCK_SIZE_BYTES 18u +#define K_BYTES_PER_THREAD 8u +#define K_BYTES_PER_INNER_LOOP 4u +#define DEQUANT_K_PACKED_TO_SHMEM dequant_q4_0_packed_to_shmem +#elif defined(K_Q8_0) +#define K_NQ 16 +#define K_BLOCK_SIZE_BYTES 34u +#define K_BYTES_PER_THREAD 16u +#define K_BYTES_PER_INNER_LOOP 4u +#define DEQUANT_K_PACKED_TO_SHMEM dequant_q8_0_packed_to_shmem +#endif + +#if defined(V_Q4_0) +#define V_NQ 16 +#define V_BLOCK_SIZE_BYTES 18u +#define V_BYTES_PER_THREAD 8u +#define V_BYTES_PER_INNER_LOOP 4u +#define DEQUANT_V_PACKED_TO_SHMEM dequant_q4_0_packed_to_shmem +#elif defined(V_Q8_0) +#define V_NQ 16 +#define V_BLOCK_SIZE_BYTES 34u +#define V_BYTES_PER_THREAD 16u +#define V_BYTES_PER_INNER_LOOP 4u +#define DEQUANT_V_PACKED_TO_SHMEM dequant_q8_0_packed_to_shmem +#endif + +#ifndef K_DIRECT +fn load_k_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, k_head_offset: u32) { +#if defined(K_Q4_0) || defined(K_Q8_0) + for (var elem_idx = local_x * K_NQ; elem_idx < kv_count * HEAD_DIM_QK; elem_idx += WG_SIZE * K_NQ) { + let blck_idx = elem_idx / BLOCK_SIZE; + let block_offset = (elem_idx % BLOCK_SIZE) / K_NQ; + let k_row = blck_idx / BLOCKS_K; + let global_k_row = kv_tile + k_row; + let block_k = blck_idx % BLOCKS_K; + let row_offset = k_row * HEAD_DIM_QK; + let global_block_idx = k_head_offset + global_k_row * params.stride_k1 + block_k; + let block_byte_base = global_block_idx * K_BLOCK_SIZE_BYTES; + let d = f16_from_u16(load_k_u16_at(block_byte_base)); + let thread_byte_offset = block_offset * K_BYTES_PER_THREAD; + let shmem_idx = row_offset + block_k * BLOCK_SIZE + thread_byte_offset; + for (var j = 0u; j < K_BYTES_PER_THREAD / K_BYTES_PER_INNER_LOOP; j += 1u) { + let q_byte_offset = block_byte_base + 2u + thread_byte_offset + j * K_BYTES_PER_INNER_LOOP; + let q_packed = load_k_u32_at(q_byte_offset); + DEQUANT_K_PACKED_TO_SHMEM(q_packed, d, shmem_idx + j * K_BYTES_PER_INNER_LOOP); + } + } +#elif defined(FLASH_ATTN_SCALAR_KV) + for (var elem_idx = local_x; elem_idx < KV_TILE * HEAD_DIM_QK; elem_idx += WG_SIZE) { + let k_row = elem_idx / HEAD_DIM_QK; + let k_col = elem_idx % HEAD_DIM_QK; + let global_k_row = kv_tile + k_row; + let global_k_row_offset = k_head_offset + global_k_row * params.stride_k1; + STAGING_SHMEM[elem_idx] = STAGING_OUT_TYPE(select( + 0.0, + K[global_k_row_offset + k_col], + global_k_row < params.seq_len_kv && k_col < HEAD_DIM_QK)); + } +#else + for (var vec_idx_local = local_x; vec_idx_local < kv_count * Q_CHUNKS; vec_idx_local += WG_SIZE) { + let kv_local = vec_idx_local / Q_CHUNKS; + let chunk = vec_idx_local % Q_CHUNKS; + let global_k_row = kv_tile + kv_local; + let k_vec_index = (k_head_offset + global_k_row * params.stride_k1 + chunk * 4u) >> 2u; + let k4 = K[k_vec_index]; + let kv_off = kv_local * HEAD_DIM_QK + chunk * 4u; + STAGING_SHMEM[kv_off + 0u] = STAGING_OUT_TYPE(k4.x); + STAGING_SHMEM[kv_off + 1u] = STAGING_OUT_TYPE(k4.y); + STAGING_SHMEM[kv_off + 2u] = STAGING_OUT_TYPE(k4.z); + STAGING_SHMEM[kv_off + 3u] = STAGING_OUT_TYPE(k4.w); + } +#endif +} +#endif // !defined(K_DIRECT) + +#ifndef V_DIRECT +fn load_v_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, v_head_offset: u32) { +#if defined(V_Q4_0) || defined(V_Q8_0) + for (var elem_idx = local_x * V_NQ; elem_idx < kv_count * HEAD_DIM_V; elem_idx += WG_SIZE * V_NQ) { + let blck_idx = elem_idx / BLOCK_SIZE; + let block_offset = (elem_idx % BLOCK_SIZE) / V_NQ; + let v_row = blck_idx / BLOCKS_V; + let global_v_row = kv_tile + v_row; + let block_k = blck_idx % BLOCKS_V; + let row_offset = v_row * HEAD_DIM_V; + let global_block_idx = v_head_offset + global_v_row * params.stride_v1 + block_k; + let block_byte_base = global_block_idx * V_BLOCK_SIZE_BYTES; + let d = f16_from_u16(load_v_u16_at(block_byte_base)); + let thread_byte_offset = block_offset * V_BYTES_PER_THREAD; + let shmem_idx = row_offset + block_k * BLOCK_SIZE + thread_byte_offset; + for (var j = 0u; j < V_BYTES_PER_THREAD / V_BYTES_PER_INNER_LOOP; j += 1u) { + let q_byte_offset = block_byte_base + 2u + thread_byte_offset + j * V_BYTES_PER_INNER_LOOP; + let q_packed = load_v_u32_at(q_byte_offset); + DEQUANT_V_PACKED_TO_SHMEM(q_packed, d, shmem_idx + j * V_BYTES_PER_INNER_LOOP); + } + } +#elif defined(FLASH_ATTN_SCALAR_KV) + for (var elem_idx = local_x; elem_idx < KV_TILE * HEAD_DIM_V; elem_idx += WG_SIZE) { + let v_row = elem_idx / HEAD_DIM_V; + let v_col = elem_idx % HEAD_DIM_V; + let global_v_row = kv_tile + v_row; + let global_v_row_offset = v_head_offset + global_v_row * params.stride_v1; + STAGING_SHMEM[elem_idx] = STAGING_OUT_TYPE(select( + 0.0, + V[global_v_row_offset + v_col], + global_v_row < params.seq_len_kv && v_col < HEAD_DIM_V)); + } +#else + for (var vec_idx_local = local_x; vec_idx_local < kv_count * V_CHUNKS; vec_idx_local += WG_SIZE) { + let kv_local = vec_idx_local / V_CHUNKS; + let chunk = vec_idx_local % V_CHUNKS; + let global_v_row = kv_tile + kv_local; + let v_vec_index = (v_head_offset + global_v_row * params.stride_v1 + chunk * 4u) >> 2u; + let v4 = V[v_vec_index]; + let kv_off = kv_local * HEAD_DIM_V + chunk * 4u; + STAGING_SHMEM[kv_off + 0u] = STAGING_OUT_TYPE(v4.x); + STAGING_SHMEM[kv_off + 1u] = STAGING_OUT_TYPE(v4.y); + STAGING_SHMEM[kv_off + 2u] = STAGING_OUT_TYPE(v4.z); + STAGING_SHMEM[kv_off + 3u] = STAGING_OUT_TYPE(v4.w); + } +#endif +} +#endif // !defined(V_DIRECT) diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_tile.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_tile.wgsl index 43f4fe7ca..8cd18b921 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_tile.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_tile.wgsl @@ -3,192 +3,32 @@ enable subgroups; #define BYTE_HELPERS #include "common_decls.tmpl" +#include "flash_attn_decls.tmpl" -#ifdef Q_F16 -#define Q_TYPE f16 -#else -#define Q_TYPE f32 -#endif - -#ifdef K_F32 -#define K_TYPE f32 -#elif defined(K_Q4_0) || defined(K_Q8_0) -#define K_TYPE u32 -#else -#define K_TYPE f16 -#endif - -#ifdef V_F32 -#define V_TYPE f32 -#elif defined(V_Q4_0) || defined(V_Q8_0) -#define V_TYPE u32 -#else -#define V_TYPE f16 -#endif - -#ifdef DST_F16 -#define DST_TYPE f16 -#else -#define DST_TYPE f32 -#endif - +// Default values +// The actual values are defined in shader-lib. #define HEAD_DIM_QK 64 #define HEAD_DIM_V 64 #define Q_TILE 4 #define KV_TILE 64 #define WG_SIZE 128 -#ifndef MIN_SUBGROUP_SIZE -#define MIN_SUBGROUP_SIZE MAX_SUBGROUP_SIZE -#endif -struct Params { - offset_q: u32, - offset_k: u32, - offset_v: u32, - offset_mask: u32, - offset_sinks: u32, - offset_dst: u32, - - n_heads: u32, - seq_len_q: u32, - seq_len_kv: u32, - - stride_q1: u32, - stride_q2: u32, - stride_q3: u32, - stride_k1: u32, - stride_k2: u32, - stride_k3: u32, - stride_v1: u32, - stride_v2: u32, - stride_v3: u32, - stride_mask3: u32, - - q_per_kv: u32, - - scale: f32, - max_bias: f32, - logit_softcap: f32, - n_head_log2: f32, - m0: f32, - m1: f32, -}; - -@group(0) @binding(0) var Q: array; -#ifdef KV_OVERLAP -#if defined(K_Q4_0) || defined(K_Q8_0) -@group(0) @binding(1) var K: array; -#else -@group(0) @binding(1) var K: array>; -#endif -#define V K -#else -#if defined(K_Q4_0) || defined(K_Q8_0) -@group(0) @binding(1) var K: array; -#else -@group(0) @binding(1) var K: array>; -#endif -#if defined(V_Q4_0) || defined(V_Q8_0) -@group(0) @binding(2) var V: array; -#else -@group(0) @binding(2) var V: array>; -#endif -#endif - -#if defined(MASK) && defined(SINKS) -#ifdef KV_OVERLAP -@group(0) @binding(2) var mask: array; -@group(0) @binding(3) var sinks: array; -#define DST_BINDING 4 -#define PARAMS_BINDING 5 -#else -@group(0) @binding(3) var mask: array; -@group(0) @binding(4) var sinks: array; -#define DST_BINDING 5 -#define PARAMS_BINDING 6 -#endif -#elif defined(MASK) -#ifdef KV_OVERLAP -@group(0) @binding(2) var mask: array; -#define DST_BINDING 3 -#define PARAMS_BINDING 4 -#else -@group(0) @binding(3) var mask: array; -#define DST_BINDING 4 -#define PARAMS_BINDING 5 -#endif -#elif defined(SINKS) -#ifdef KV_OVERLAP -@group(0) @binding(2) var sinks: array; -#define DST_BINDING 3 -#define PARAMS_BINDING 4 -#else -@group(0) @binding(3) var sinks: array; -#define DST_BINDING 4 -#define PARAMS_BINDING 5 -#endif -#else -#ifdef KV_OVERLAP -#define DST_BINDING 2 -#define PARAMS_BINDING 3 -#else -#define DST_BINDING 3 -#define PARAMS_BINDING 4 -#endif -#endif - -@group(0) @binding(DST_BINDING) var dst: array>; -@group(0) @binding(PARAMS_BINDING) var params: Params; - -const FLOAT_MIN: f32 = -1.0e9; const Q_CHUNKS: u32 = HEAD_DIM_QK / 4u; const V_CHUNKS: u32 = HEAD_DIM_V / 4u; const SCORE_REGS_PER_LANE: u32 = (KV_TILE + MIN_SUBGROUP_SIZE - 1u) / MIN_SUBGROUP_SIZE; const OUT_REGS_PER_LANE: u32 = (V_CHUNKS + MIN_SUBGROUP_SIZE - 1u) / MIN_SUBGROUP_SIZE; + +#if !defined(K_DIRECT) || !defined(V_DIRECT) +#define STAGING_SHMEM kv_shmem +#define STAGING_OUT_TYPE f16 +#include "flash_attn_staging.tmpl" const kv_shmem_size = KV_TILE * max(HEAD_DIM_QK, HEAD_DIM_V); +var kv_shmem: array; +#endif var q_shmem: array; -var kv_shmem: array; var p_shmem: array; -#define QUANT_SHMEM kv_shmem -#define QUANT_OUT_TYPE f16 -#include "flash_attn_quant_staging.tmpl" - -#if !defined(K_Q4_0) && !defined(K_Q8_0) -fn load_k_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, k_head_offset: u32) { - for (var vec_idx_local = local_x; vec_idx_local < kv_count * Q_CHUNKS; vec_idx_local += WG_SIZE) { - let kv_local = vec_idx_local / Q_CHUNKS; - let chunk = vec_idx_local % Q_CHUNKS; - let global_k_row = kv_tile + kv_local; - let k_vec_index = (k_head_offset + global_k_row * params.stride_k1 + chunk * 4u) >> 2u; - let k4 = K[k_vec_index]; - let kv_off = kv_local * HEAD_DIM_QK + chunk * 4u; - kv_shmem[kv_off + 0u] = f16(k4.x); - kv_shmem[kv_off + 1u] = f16(k4.y); - kv_shmem[kv_off + 2u] = f16(k4.z); - kv_shmem[kv_off + 3u] = f16(k4.w); - } -} -#endif - -#if !defined(V_Q4_0) && !defined(V_Q8_0) -fn load_v_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, v_head_offset: u32) { - for (var vec_idx_local = local_x; vec_idx_local < kv_count * V_CHUNKS; vec_idx_local += WG_SIZE) { - let kv_local = vec_idx_local / V_CHUNKS; - let chunk = vec_idx_local % V_CHUNKS; - let global_v_row = kv_tile + kv_local; - let v_vec_index = (v_head_offset + global_v_row * params.stride_v1 + chunk * 4u) >> 2u; - let v4 = V[v_vec_index]; - let kv_off = kv_local * HEAD_DIM_V + chunk * 4u; - kv_shmem[kv_off + 0u] = f16(v4.x); - kv_shmem[kv_off + 1u] = f16(v4.y); - kv_shmem[kv_off + 2u] = f16(v4.z); - kv_shmem[kv_off + 3u] = f16(v4.w); - } -} -#endif - @compute @workgroup_size(WG_SIZE) fn main(@builtin(workgroup_id) wg_id: vec3, @builtin(local_invocation_id) local_id: vec3, diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_split.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_split.wgsl index b8e0be90d..42f3b1089 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_split.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_split.wgsl @@ -4,200 +4,35 @@ enable subgroups; #define BYTE_HELPERS #include "common_decls.tmpl" +#define FLASH_ATTN_VEC_SPLIT +#include "flash_attn_decls.tmpl" -#ifdef K_F32 -#define K_TYPE f32 -#elif defined(K_Q4_0) || defined(K_Q8_0) -#define K_TYPE u32 -#else -#define K_TYPE f16 -#endif - -#ifdef V_F32 -#define V_TYPE f32 -#elif defined(V_Q4_0) || defined(V_Q8_0) -#define V_TYPE u32 -#else -#define V_TYPE f16 -#endif - -#ifdef Q_F16 -#define Q_TYPE f16 -#else -#define Q_TYPE f32 -#endif - -#ifdef DST_F16 -#define DST_TYPE f16 -#else -#define DST_TYPE f32 -#endif - +// Default values +// The actual values are defined in shader-lib. #define HEAD_DIM_QK 64 #define HEAD_DIM_V 64 - -#define KV_GRANULARITY 8 #define KV_TILE 16 #define WG_SIZE 64 -#define KV_BLOCKS (KV_TILE / KV_GRANULARITY) - -struct Params { - offset_q: u32, - offset_k: u32, - offset_v: u32, - offset_mask: u32, - offset_sinks: u32, - offset_dst: u32, - - // shapes of Q/K/V - n_heads: u32, - seq_len_q: u32, - seq_len_kv: u32, - - // strides (in elements) - stride_q1: u32, - stride_q2: u32, - stride_q3: u32, - stride_k1: u32, - stride_k2: u32, - stride_k3: u32, - stride_v1: u32, - stride_v2: u32, - stride_v3: u32, - stride_mask3: u32, - - // repeat factors for K/V, e.g., MHA vs. MQA vs. GQA - q_per_kv: u32, - - // softmax params - scale: f32, - max_bias: f32, - logit_softcap: f32, - n_head_log2: f32, - m0: f32, - m1: f32, - -#ifdef BLK - blk_base: u32, - blk_nblk0: u32, - blk_nblk1: u32, -#endif - - tmp_data_base: u32, - tmp_stats_base: u32, - nwg: u32, -}; - -@group(0) @binding(0) var Q: array; -#ifdef KV_OVERLAP -#if defined(K_Q4_0) || defined(K_Q8_0) -@group(0) @binding(1) var K: array; -#else -@group(0) @binding(1) var K: array>; -#endif -#define V K -#else -#if defined(K_Q4_0) || defined(K_Q8_0) -@group(0) @binding(1) var K: array; -#else -@group(0) @binding(1) var K: array>; -#endif -#if defined(V_Q4_0) || defined(V_Q8_0) -@group(0) @binding(2) var V: array; -#else -@group(0) @binding(2) var V: array>; -#endif -#endif -#if defined(MASK) && defined(SINKS) -#ifdef KV_OVERLAP -@group(0) @binding(2) var mask: array; -@group(0) @binding(3) var sinks: array; -#ifdef BLK -#define BLK_BINDING 4 -#define TMP_BINDING 5 -#define DST_BINDING 6 -#define PARAMS_BINDING 7 -#else -#define TMP_BINDING 4 -#define DST_BINDING 5 -#define PARAMS_BINDING 6 -#endif -#else -@group(0) @binding(3) var mask: array; -@group(0) @binding(4) var sinks: array; -#ifdef BLK -#define BLK_BINDING 5 -#define TMP_BINDING 6 -#define DST_BINDING 7 -#define PARAMS_BINDING 8 -#else -#define TMP_BINDING 5 -#define DST_BINDING 6 -#define PARAMS_BINDING 7 -#endif -#endif -#elif defined(MASK) -#ifdef KV_OVERLAP -@group(0) @binding(2) var mask: array; -#ifdef BLK -#define BLK_BINDING 3 -#define TMP_BINDING 4 -#define DST_BINDING 5 -#define PARAMS_BINDING 6 -#else -#define TMP_BINDING 3 -#define DST_BINDING 4 -#define PARAMS_BINDING 5 -#endif -#else -@group(0) @binding(3) var mask: array; -#ifdef BLK -#define BLK_BINDING 4 -#define TMP_BINDING 5 -#define DST_BINDING 6 -#define PARAMS_BINDING 7 -#else -#define TMP_BINDING 4 -#define DST_BINDING 5 -#define PARAMS_BINDING 6 -#endif -#endif -#elif defined(SINKS) -#ifdef KV_OVERLAP -@group(0) @binding(2) var sinks: array; -#define TMP_BINDING 3 -#define DST_BINDING 4 -#define PARAMS_BINDING 5 -#else -@group(0) @binding(3) var sinks: array; -#define TMP_BINDING 4 -#define DST_BINDING 5 -#define PARAMS_BINDING 6 -#endif -#else -#ifdef KV_OVERLAP -#define TMP_BINDING 2 -#define DST_BINDING 3 -#define PARAMS_BINDING 4 -#else -#define TMP_BINDING 3 -#define DST_BINDING 4 -#define PARAMS_BINDING 5 -#endif -#endif - -#ifdef BLK -@group(0) @binding(BLK_BINDING) var blk: array; -#endif -@group(0) @binding(TMP_BINDING) var tmp: array; -@group(0) @binding(DST_BINDING) var dst: array>; -@group(0) @binding(PARAMS_BINDING) var params: Params; - -// Just a very small float value. -const FLOAT_MIN: f32 = -1.0e9; +const Q_CHUNKS: u32 = HEAD_DIM_QK / 4u; +const V_CHUNKS: u32 = HEAD_DIM_V / 4u; const kv_shmem_size = KV_TILE * max(HEAD_DIM_QK, HEAD_DIM_V); +#if defined(K_DIRECT) || defined(V_DIRECT) +// Shared memory for scale factor (d) in quantized K/V. Multiple threads use the same value, +// so caching it is more efficient, even on the direct path. +var d_shmem: array; +#endif + +// K/V shared memory handling +#if !defined(K_DIRECT) || !defined(V_DIRECT) +#define STAGING_SHMEM kv_shmem +#define STAGING_OUT_TYPE f32 +#include "flash_attn_staging.tmpl" +// we can reuse the same shmem for K and V since we only need one at a time +var kv_shmem: array; +#endif + var q_shmem: array; var o_shmem: array; // note that we reuse the same storage for both since we only need one at a time @@ -208,59 +43,6 @@ var inter_shmem: array; var mask_shmem: array; #endif -#if defined(K_DIRECT) || defined(V_DIRECT) -// Shared memory for scale factor (d) in quantized K/V. Multiple threads use the same value, -// so caching it is more efficient, even on the direct path. -var d_shmem: array; -#endif - -// K/V shared memory handling -#if !defined(K_DIRECT) || !defined(V_DIRECT) - -// we can reuse the same shmem for K and V since we only need one at a time -var kv_shmem: array; - -#define QUANT_SHMEM kv_shmem -#define QUANT_OUT_TYPE f32 -#include "flash_attn_quant_staging.tmpl" - -#if !defined(K_DIRECT) && !defined(K_Q4_0) && !defined(K_Q8_0) -fn load_k_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, k_head_offset: u32) { - for (var elem_idx = local_x * 4u; elem_idx < KV_TILE * HEAD_DIM_QK; elem_idx += WG_SIZE * 4u) { - let k_row = elem_idx / HEAD_DIM_QK; - let k_col = elem_idx % HEAD_DIM_QK; - let global_k_row = kv_tile + k_row; - let global_k_row_offset = k_head_offset + global_k_row * params.stride_k1; - let in_bounds = global_k_row < params.seq_len_kv && (k_col + 3u) < HEAD_DIM_QK; - let vec_idx = (global_k_row_offset + k_col) >> 2u; - let k4 = select(vec4(0.0), K[vec_idx], in_bounds); - kv_shmem[elem_idx + 0u] = f32(k4.x); - kv_shmem[elem_idx + 1u] = f32(k4.y); - kv_shmem[elem_idx + 2u] = f32(k4.z); - kv_shmem[elem_idx + 3u] = f32(k4.w); - } -} -#endif - -#if !defined(V_DIRECT) && !defined(V_Q4_0) && !defined(V_Q8_0) -fn load_v_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, v_head_offset: u32) { - for (var elem_idx = local_x * 4u; elem_idx < KV_TILE * HEAD_DIM_V; elem_idx += WG_SIZE * 4u) { - let v_row = elem_idx / HEAD_DIM_V; - let v_col = elem_idx % HEAD_DIM_V; - let global_v_row = kv_tile + v_row; - let global_v_row_offset = v_head_offset + global_v_row * params.stride_v1; - let in_bounds = global_v_row < params.seq_len_kv && (v_col + 3u) < HEAD_DIM_V; - let vec_idx = (global_v_row_offset + v_col) >> 2u; - let v4 = select(vec4(0.0), V[vec_idx], in_bounds); - kv_shmem[elem_idx + 0u] = f32(v4.x); - kv_shmem[elem_idx + 1u] = f32(v4.y); - kv_shmem[elem_idx + 2u] = f32(v4.z); - kv_shmem[elem_idx + 3u] = f32(v4.w); - } -} -#endif -#endif // !defined(K_DIRECT) || !defined(V_DIRECT) - // Storage for row max and exp sum during online softmax fn calc_softmax_term(kv_idx: u32, slope: f32, has_bias: bool, apply_mask: bool) -> f32 { var v = select(FLOAT_MIN, diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/im2col.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/im2col.wgsl index 386ebab87..ebcf031c3 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/im2col.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/im2col.wgsl @@ -1,19 +1,9 @@ -#include "common_decls.tmpl" enable f16; @group(0) @binding(0) -#if defined(INPUT_F32) -var input: array; -#elif defined(INPUT_F16) -var input: array; -#endif - +var input: array; @group(0) @binding(1) -#if defined(OUTPUT_F32) -var output: array; -#elif defined(OUTPUT_F16) -var output: array; -#endif +var output: array; struct Params { offset_i: u32, @@ -38,22 +28,6 @@ struct Params { @group(0) @binding(2) var params: Params; -fn load_input(idx: u32) -> f32 { - #if defined(INPUT_F32) - return input[idx]; - #elif defined(INPUT_F16) - return f32(input[idx]); - #endif -} - -fn store_output(idx: u32, val: f32) { - #if defined(OUTPUT_F32) - output[idx] = val; - #elif defined(OUTPUT_F16) - output[idx] = f16(val); - #endif -} - @compute @workgroup_size(WG_SIZE) fn main( @builtin(global_invocation_id) gid: vec3, @@ -90,12 +64,14 @@ fn main( let iw_i32 = i32(ow * params.s0 + kw * params.d0) - i32(params.p0); let ih_i32 = i32(oh * params.s1 + kh * params.d1) - i32(params.p1); + let output_idx = params.offset_o + k * params.so0 + ow * params.so1 + oh * params.so2 + n * params.so3; + if (iw_i32 >= 0 && iw_i32 < i32(params.IW) && ih_i32 >= 0 && ih_i32 < i32(params.IH)) { let iw = u32(iw_i32); let ih = u32(ih_i32); let in_idx = params.offset_i + iw * params.si0 + ih * params.si1 + ic * params.si2 + n * params.si3; - store_output(params.offset_o + k * params.so0 + ow * params.so1 + oh * params.so2 + n * params.so3, load_input(in_idx)); + output[output_idx] = OUTPUT_TYPE(input[in_idx]); } else { - store_output(params.offset_o + k * params.so0 + ow * params.so1 + oh * params.so2 + n * params.so3, 0.0); + output[output_idx] = OUTPUT_TYPE(0.0); } } diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/rms_norm_mul.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/rms_norm_mul.wgsl index fd20a4e54..c9e424ffc 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/rms_norm_mul.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/rms_norm_mul.wgsl @@ -88,7 +88,6 @@ struct Params { ne0: u32, ne1: u32, ne2: u32, - ne3: u32, eps: f32 }; diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/row_norm.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/row_norm.wgsl index 5eaf5e7bb..7629bf5b4 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/row_norm.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/row_norm.wgsl @@ -31,7 +31,6 @@ struct Params { ne0: u32, ne1: u32, ne2: u32, - ne3: u32, eps: f32 }; diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/soft_max.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/soft_max.wgsl index 10edf1360..1c29a9221 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/soft_max.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/soft_max.wgsl @@ -27,7 +27,6 @@ struct Params { stride_dst3: u32, // shape of src0/dst - ne: u32, ne0: u32, ne1: u32, ne2: u32, @@ -43,71 +42,38 @@ struct Params { m1: f32, }; -@group(0) @binding(0) +#define SRC_BINDING 0 +@group(0) @binding(SRC_BINDING) var src: array; #ifdef HAS_MASK -#ifdef HAS_SINK -@group(0) @binding(1) +#define MASK_BINDING SRC_BINDING + 1 +@group(0) @binding(MASK_BINDING) var mask: array; -@group(0) @binding(2) -var sinks: array; - -#ifdef INPLACE -@group(0) @binding(3) -var params: Params; - #else -@group(0) @binding(3) -var dst: array; -@group(0) @binding(4) -var params: Params; +#define MASK_BINDING SRC_BINDING #endif -#else -@group(0) @binding(1) -var mask: array; - -#ifdef INPLACE -@group(0) @binding(2) -var params: Params; - -#else -@group(0) @binding(2) -var dst: array; -@group(0) @binding(3) -var params: Params; -#endif -#endif - -#else #ifdef HAS_SINK -@group(0) @binding(1) +#define SINKS_BINDING MASK_BINDING + 1 +@group(0) @binding(SINKS_BINDING) var sinks: array; +#else +#define SINKS_BINDING MASK_BINDING +#endif + +#define DST_BINDING SINKS_BINDING + 1 +@group(0) @binding(DST_BINDING) +var dst: array; #ifdef INPLACE -@group(0) @binding(2) -var params: Params; - +#define PARAMS_BINDING DST_BINDING #else -@group(0) @binding(2) -var dst: array; -@group(0) @binding(3) -var params: Params; +#define PARAMS_BINDING (DST_BINDING + 1) #endif -#else -#ifdef INPLACE -@group(0) @binding(1) +@group(0) @binding(PARAMS_BINDING) var params: Params; -#else -@group(0) @binding(1) -var dst: array; -@group(0) @binding(2) -var params: Params; -#endif -#endif -#endif #ifdef INPLACE fn inter_value(i: u32) -> f32 { @@ -242,4 +208,3 @@ fn main(@builtin(workgroup_id) wid: vec3, col += WG_SIZE; } } - diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/solve_tri.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/solve_tri.wgsl index 9d5d902cb..c01df92f0 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/solve_tri.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/solve_tri.wgsl @@ -29,7 +29,6 @@ struct Params { k: u32, ne2: u32, - ne3: u32, }; @group(0) @binding(3) diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/ssm_scan.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/ssm_scan.wgsl index 66bfdd640..2d4c4e5a0 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/ssm_scan.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/ssm_scan.wgsl @@ -39,7 +39,6 @@ struct Params { n_head: u32, n_group: u32, n_seq_tokens: u32, - n_seqs: u32, y_elems: u32, };