]> git.djapps.eu Git - pkg/ggml/sources/llama.cpp/commitdiff
ggml-webgpu : refactor several wgsl files and simplify flash_attn wgsl. (#26134)
authorMasashi Yoshimura <redacted>
Mon, 10 Aug 2026 06:29:41 +0000 (15:29 +0900)
committerGitHub <redacted>
Mon, 10 Aug 2026 06:29:41 +0000 (09:29 +0300)
17 files changed:
ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp
ggml/src/ggml-webgpu/ggml-webgpu.cpp
ggml/src/ggml-webgpu/wgsl-shaders/concat.wgsl
ggml/src/ggml-webgpu/wgsl-shaders/conv2d.wgsl
ggml/src/ggml-webgpu/wgsl-shaders/conv2d_dw.wgsl
ggml/src/ggml-webgpu/wgsl-shaders/flash_attn.wgsl
ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_decls.tmpl [new file with mode: 0644]
ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_quant_staging.tmpl [deleted file]
ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_staging.tmpl [new file with mode: 0644]
ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_tile.wgsl
ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_vec_split.wgsl
ggml/src/ggml-webgpu/wgsl-shaders/im2col.wgsl
ggml/src/ggml-webgpu/wgsl-shaders/rms_norm_mul.wgsl
ggml/src/ggml-webgpu/wgsl-shaders/row_norm.wgsl
ggml/src/ggml-webgpu/wgsl-shaders/soft_max.wgsl
ggml/src/ggml-webgpu/wgsl-shaders/solve_tri.wgsl
ggml/src/ggml-webgpu/wgsl-shaders/ssm_scan.wgsl

index 66c1c3c8977e38c61286962d0defa2ff4ac28c9c..35a55ecaf64409c0f1ea8cf6101ca6fd08983dfa 100644 (file)
@@ -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));
 
index c001cda7d116c8a34d76de1281c292e222bcc402..ba4b91695faea65a164c8db1351b58f473a12af9 100644 (file)
@@ -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<wgpu::BindGroupEntry> 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<webgpu_encoded_op> 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],
index eb901bf054717096e2201731ad3947c6391e41d1..7ccad73f4b3ae054b2c95c75c61771148f9a598f 100644 (file)
@@ -18,7 +18,6 @@ struct Params {
     ne0: u32,
     ne1: u32,
     ne2: u32,
-    ne3: u32,
 
     dim: u32,
     src0_nedim: u32
index 9eb131dc22184e346efc9580bd4ddf9629ac726c..38c714ba5990be725591462c7c1cd6a9a9569d45 100644 (file)
@@ -2,25 +2,11 @@
 enable f16;
 
 @group(0) @binding(0)
-#if defined(WEIGHT_F32)
-var<storage, read_write> weights: array<f32>;
-#elif defined(WEIGHT_F16)
-var<storage, read_write> weights: array<f16>;
-#endif
-
+var<storage, read_write> weights: array<WEIGHT_TYPE>;
 @group(0) @binding(1)
-#if defined(INPUT_F32)
-var<storage, read_write> input: array<f32>;
-#elif defined(INPUT_F16)
-var<storage, read_write> input: array<f16>;
-#endif
-
+var<storage, read_write> input: array<INPUT_TYPE>;
 @group(0) @binding(2)
-#if defined(OUTPUT_F32)
-var<storage, read_write> output: array<f32>;
-#elif defined(OUTPUT_F16)
-var<storage, read_write> output: array<f16>;
-#endif
+var<storage, read_write> output: array<OUTPUT_TYPE>;
 
 struct Params {
     offset_w: u32,
@@ -50,30 +36,6 @@ struct Params {
 @group(0) @binding(3)
 var<uniform> 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);
 }
index 42d6f027cab6b9892f060ac495130a420d9d13b8..fc028e42998bfe3608df03d830a8b9aa6355c0be 100644 (file)
@@ -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<storage, read_write> weights: array<f32>;
-#elif defined(WEIGHT_F16)
-var<storage, read_write> weights: array<f16>;
-#endif
-
+var<storage, read_write> weights: array<WEIGHT_TYPE>;
 @group(0) @binding(1)
-#if defined(INPUT_F32)
-var<storage, read_write> input: array<f32>;
-#elif defined(INPUT_F16)
-var<storage, read_write> input: array<f16>;
-#endif
-
+var<storage, read_write> input: array<INPUT_TYPE>;
 @group(0) @binding(2)
-#if defined(OUTPUT_F32)
-var<storage, read_write> output: array<f32>;
-#elif defined(OUTPUT_F16)
-var<storage, read_write> output: array<f16>;
-#endif
+var<storage, read_write> output: array<OUTPUT_TYPE>;
 
 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<uniform> 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));
 }
index 75f33e68ae53f693746a0967bf72a4f999168b58..d5bf2af8d2c862ee1db354b37de830fd994a7642 100644 (file)
@@ -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<storage, read_write> Q: array<f32>;
-#ifdef KV_OVERLAP
-@group(0) @binding(1) var<storage, read_write> K: array<K_TYPE>;
-#define V K
-#else
-@group(0) @binding(1) var<storage, read_write> K: array<K_TYPE>;
-@group(0) @binding(2) var<storage, read_write> V: array<V_TYPE>;
-#endif
-
-#if defined(MASK) && defined(SINKS)
-#ifdef KV_OVERLAP
-@group(0) @binding(2) var<storage, read_write> mask: array<f16>;
-@group(0) @binding(3) var<storage, read_write> sinks: array<f32>;
-#define DST_BINDING 4
-#define PARAMS_BINDING 5
-#else
-@group(0) @binding(3) var<storage, read_write> mask: array<f16>;
-@group(0) @binding(4) var<storage, read_write> sinks: array<f32>;
-#define DST_BINDING 5
-#define PARAMS_BINDING 6
-#endif
-#elif defined(MASK)
-#ifdef KV_OVERLAP
-@group(0) @binding(2) var<storage, read_write> mask: array<f16>;
-#define DST_BINDING 3
-#define PARAMS_BINDING 4
-#else
-@group(0) @binding(3) var<storage, read_write> mask: array<f16>;
-#define DST_BINDING 4
-#define PARAMS_BINDING 5
-#endif
-#elif defined(SINKS)
-#ifdef KV_OVERLAP
-@group(0) @binding(2) var<storage, read_write> sinks: array<f32>;
-#define DST_BINDING 3
-#define PARAMS_BINDING 4
-#else
-@group(0) @binding(3) var<storage, read_write> sinks: array<f32>;
-#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<storage, read_write> dst: array<vec4<f32>>;
-@group(0) @binding(PARAMS_BINDING) var<uniform> params: Params;
-
-// Just a very small float value.
-const FLOAT_MIN: f32 = -1.0e9;
-
 // The number of Q rows processed per workgroup
 var<workgroup> q_shmem: array<f16, Q_TILE * HEAD_DIM_QK>;
 
 #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<workgroup> kv_shmem: array<f16, kv_shmem_size>;
@@ -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<storage, array<vec4<f32>>, read_write>, scalar_index: u32) -> vec4<f32> {
-    return (*buf)[scalar_index >> 2u];
-}
-
-fn load_kx4(buf: ptr<storage, array<vec4<K_TYPE>>, read_write>, scalar_index: u32) -> vec4<K_TYPE> {
-    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<u32>,
     @builtin(local_invocation_id) local_id: vec3<u32>,
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 (file)
index 0000000..48a79b6
--- /dev/null
@@ -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<K_TYPE>
+#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<V_TYPE>
+#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<storage, read_write> Q: array<Q_TYPE>;
+@group(0) @binding(1) var<storage, read_write> K: array<K_STORAGE_TYPE>;
+#ifdef KV_OVERLAP
+#define V K
+#define MASK_BINDING 2
+#else
+@group(0) @binding(2) var<storage, read_write> V: array<V_STORAGE_TYPE>;
+#define MASK_BINDING 3
+#endif // KV_OVERLAP
+
+#ifdef MASK
+@group(0) @binding(MASK_BINDING) var<storage, read_write> mask: array<f16>;
+#define SINKS_BINDING (MASK_BINDING + 1)
+#else
+#define SINKS_BINDING MASK_BINDING
+#endif
+
+#ifdef SINKS
+@group(0) @binding(SINKS_BINDING) var<storage, read_write> sinks: array<f32>;
+#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<storage, read_write> blk: array<u32>;
+#define TMP_BINDING (BLK_BINDING + 1)
+#else
+#define TMP_BINDING BLK_BINDING
+#endif
+
+@group(0) @binding(TMP_BINDING) var<storage, read_write> tmp: array<f32>;
+#define DST_BINDING (TMP_BINDING + 1)
+#else
+#define DST_BINDING BLK_BINDING
+#endif // FLASH_ATTN_VEC_SPLIT
+
+@group(0) @binding(DST_BINDING) var<storage, read_write> dst: array<vec4<DST_TYPE>>;
+
+#define PARAMS_BINDING (DST_BINDING + 1)
+@group(0) @binding(PARAMS_BINDING) var<uniform> 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 (file)
index 1c23260..0000000
+++ /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 (file)
index 0000000..457df07
--- /dev/null
@@ -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)
index 43f4fe7caccd6396d104240e0aa0e3676974273d..8cd18b92184858c8950ae353836ea3a3a6db51ee 100644 (file)
@@ -3,191 +3,31 @@ 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<storage, read_write> Q: array<Q_TYPE>;
-#ifdef KV_OVERLAP
-#if defined(K_Q4_0) || defined(K_Q8_0)
-@group(0) @binding(1) var<storage, read_write> K: array<K_TYPE>;
-#else
-@group(0) @binding(1) var<storage, read_write> K: array<vec4<K_TYPE>>;
-#endif
-#define V K
-#else
-#if defined(K_Q4_0) || defined(K_Q8_0)
-@group(0) @binding(1) var<storage, read_write> K: array<K_TYPE>;
-#else
-@group(0) @binding(1) var<storage, read_write> K: array<vec4<K_TYPE>>;
-#endif
-#if defined(V_Q4_0) || defined(V_Q8_0)
-@group(0) @binding(2) var<storage, read_write> V: array<V_TYPE>;
-#else
-@group(0) @binding(2) var<storage, read_write> V: array<vec4<V_TYPE>>;
-#endif
-#endif
-
-#if defined(MASK) && defined(SINKS)
-#ifdef KV_OVERLAP
-@group(0) @binding(2) var<storage, read_write> mask: array<f16>;
-@group(0) @binding(3) var<storage, read_write> sinks: array<f32>;
-#define DST_BINDING 4
-#define PARAMS_BINDING 5
-#else
-@group(0) @binding(3) var<storage, read_write> mask: array<f16>;
-@group(0) @binding(4) var<storage, read_write> sinks: array<f32>;
-#define DST_BINDING 5
-#define PARAMS_BINDING 6
-#endif
-#elif defined(MASK)
-#ifdef KV_OVERLAP
-@group(0) @binding(2) var<storage, read_write> mask: array<f16>;
-#define DST_BINDING 3
-#define PARAMS_BINDING 4
-#else
-@group(0) @binding(3) var<storage, read_write> mask: array<f16>;
-#define DST_BINDING 4
-#define PARAMS_BINDING 5
-#endif
-#elif defined(SINKS)
-#ifdef KV_OVERLAP
-@group(0) @binding(2) var<storage, read_write> sinks: array<f32>;
-#define DST_BINDING 3
-#define PARAMS_BINDING 4
-#else
-@group(0) @binding(3) var<storage, read_write> sinks: array<f32>;
-#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<storage, read_write> dst: array<vec4<DST_TYPE>>;
-@group(0) @binding(PARAMS_BINDING) var<uniform> 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;
-const kv_shmem_size = KV_TILE * max(HEAD_DIM_QK, HEAD_DIM_V);
 
-var<workgroup> q_shmem: array<Q_TYPE, Q_TILE * HEAD_DIM_QK>;
+#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<workgroup> kv_shmem: array<f16, kv_shmem_size>;
-var<workgroup> p_shmem: array<f16, Q_TILE * KV_TILE>;
-
-#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
+var<workgroup> q_shmem: array<Q_TYPE, Q_TILE * HEAD_DIM_QK>;
+var<workgroup> p_shmem: array<f16, Q_TILE * KV_TILE>;
 
 @compute @workgroup_size(WG_SIZE)
 fn main(@builtin(workgroup_id) wg_id: vec3<u32>,
index b8e0be90d998ad56dd10df2dcb43753bef775989..42f3b10890560ea0a243ec61111bb9b3409cf36c 100644 (file)
@@ -4,210 +4,20 @@ 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<storage, read_write> Q: array<Q_TYPE>;
-#ifdef KV_OVERLAP
-#if defined(K_Q4_0) || defined(K_Q8_0)
-@group(0) @binding(1) var<storage, read_write> K: array<K_TYPE>;
-#else
-@group(0) @binding(1) var<storage, read_write> K: array<vec4<K_TYPE>>;
-#endif
-#define V K
-#else
-#if defined(K_Q4_0) || defined(K_Q8_0)
-@group(0) @binding(1) var<storage, read_write> K: array<K_TYPE>;
-#else
-@group(0) @binding(1) var<storage, read_write> K: array<vec4<K_TYPE>>;
-#endif
-#if defined(V_Q4_0) || defined(V_Q8_0)
-@group(0) @binding(2) var<storage, read_write> V: array<V_TYPE>;
-#else
-@group(0) @binding(2) var<storage, read_write> V: array<vec4<V_TYPE>>;
-#endif
-#endif
-#if defined(MASK) && defined(SINKS)
-#ifdef KV_OVERLAP
-@group(0) @binding(2) var<storage, read_write> mask: array<f16>;
-@group(0) @binding(3) var<storage, read_write> sinks: array<f32>;
-#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<storage, read_write> mask: array<f16>;
-@group(0) @binding(4) var<storage, read_write> sinks: array<f32>;
-#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<storage, read_write> mask: array<f16>;
-#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<storage, read_write> mask: array<f16>;
-#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<storage, read_write> sinks: array<f32>;
-#define TMP_BINDING 3
-#define DST_BINDING 4
-#define PARAMS_BINDING 5
-#else
-@group(0) @binding(3) var<storage, read_write> sinks: array<f32>;
-#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<storage, read_write> blk: array<u32>;
-#endif
-@group(0) @binding(TMP_BINDING) var<storage, read_write> tmp: array<f32>;
-@group(0) @binding(DST_BINDING) var<storage, read_write> dst: array<vec4<DST_TYPE>>;
-@group(0) @binding(PARAMS_BINDING) var<uniform> 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);
 
-var<workgroup> q_shmem: array<f32, HEAD_DIM_QK>;
-var<workgroup> o_shmem: array<f32, HEAD_DIM_V>;
-// note that we reuse the same storage for both since we only need one at a time
-var<workgroup> inter_shmem: array<f32, KV_TILE>;
-
-#ifdef MASK
-// storage for mask values
-var<workgroup> mask_shmem: array<f32, KV_TILE>;
-#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.
@@ -216,50 +26,22 @@ var<workgroup> d_shmem: array<f32, kv_shmem_size / 32>;
 
 // 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<workgroup> kv_shmem: array<f32, kv_shmem_size>;
-
-#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<K_TYPE>(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<V_TYPE>(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);
-    }
-}
+var<workgroup> q_shmem: array<f32, HEAD_DIM_QK>;
+var<workgroup> o_shmem: array<f32, HEAD_DIM_V>;
+// note that we reuse the same storage for both since we only need one at a time
+var<workgroup> inter_shmem: array<f32, KV_TILE>;
+
+#ifdef MASK
+// storage for mask values
+var<workgroup> mask_shmem: array<f32, KV_TILE>;
 #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 {
index 386ebab879fbd08b47e1ceb2cf18930402dbb30c..ebcf031c3bb93ccad5349921489b9e6f58798c4b 100644 (file)
@@ -1,19 +1,9 @@
-#include "common_decls.tmpl"
 enable f16;
 
 @group(0) @binding(0)
-#if defined(INPUT_F32)
-var<storage, read_write> input: array<f32>;
-#elif defined(INPUT_F16)
-var<storage, read_write> input: array<f16>;
-#endif
-
+var<storage, read_write> input: array<INPUT_TYPE>;
 @group(0) @binding(1)
-#if defined(OUTPUT_F32)
-var<storage, read_write> output: array<f32>;
-#elif defined(OUTPUT_F16)
-var<storage, read_write> output: array<f16>;
-#endif
+var<storage, read_write> output: array<OUTPUT_TYPE>;
 
 struct Params {
     offset_i: u32,
@@ -38,22 +28,6 @@ struct Params {
 @group(0) @binding(2)
 var<uniform> 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<u32>,
@@ -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);
     }
 }
index fd20a4e54c932a616f5e054acb2d1b354493c266..c9e424ffce878068174855c11f3071abd7f7060f 100644 (file)
@@ -88,7 +88,6 @@ struct Params {
     ne0: u32,
     ne1: u32,
     ne2: u32,
-    ne3: u32,
 
     eps: f32
 };
index 5eaf5e7bbe5de0af86103e67118926d7189cc810..7629bf5b4573bc56384143ac47fc378844751a75 100644 (file)
@@ -31,7 +31,6 @@ struct Params {
     ne0: u32,
     ne1: u32,
     ne2: u32,
-    ne3: u32,
 
     eps: f32
 };
index 10edf1360489f4a50516be2a52507bf89012b8ea..1c29a9221b684cce75cfcef2e57aab97c9fcee20 100644 (file)
@@ -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<storage, read_write> src: array<f32>;
 
 #ifdef HAS_MASK
-#ifdef HAS_SINK
-@group(0) @binding(1)
+#define MASK_BINDING SRC_BINDING + 1
+@group(0) @binding(MASK_BINDING)
 var<storage, read_write> mask: array<MaskType>;
-@group(0) @binding(2)
-var<storage, read_write> sinks: array<f32>;
-
-#ifdef INPLACE
-@group(0) @binding(3)
-var<uniform> params: Params;
-
 #else
-@group(0) @binding(3)
-var<storage, read_write> dst: array<f32>;
-@group(0) @binding(4)
-var<uniform> params: Params;
+#define MASK_BINDING SRC_BINDING
 #endif
 
+#ifdef HAS_SINK
+#define SINKS_BINDING MASK_BINDING + 1
+@group(0) @binding(SINKS_BINDING)
+var<storage, read_write> sinks: array<f32>;
 #else
-@group(0) @binding(1)
-var<storage, read_write> mask: array<MaskType>;
-
-#ifdef INPLACE
-@group(0) @binding(2)
-var<uniform> params: Params;
-
-#else
-@group(0) @binding(2)
-var<storage, read_write> dst: array<f32>;
-@group(0) @binding(3)
-var<uniform> params: Params;
-#endif
+#define SINKS_BINDING MASK_BINDING
 #endif
 
-#else
-#ifdef HAS_SINK
-@group(0) @binding(1)
-var<storage, read_write> sinks: array<f32>;
+#define DST_BINDING SINKS_BINDING + 1
+@group(0) @binding(DST_BINDING)
+var<storage, read_write> dst: array<f32>;
 
 #ifdef INPLACE
-@group(0) @binding(2)
-var<uniform> params: Params;
-
+#define PARAMS_BINDING DST_BINDING
 #else
-@group(0) @binding(2)
-var<storage, read_write> dst: array<f32>;
-@group(0) @binding(3)
-var<uniform> params: Params;
+#define PARAMS_BINDING (DST_BINDING + 1)
 #endif
 
-#else
-#ifdef INPLACE
-@group(0) @binding(1)
-var<uniform> params: Params;
-#else
-@group(0) @binding(1)
-var<storage, read_write> dst: array<f32>;
-@group(0) @binding(2)
+@group(0) @binding(PARAMS_BINDING)
 var<uniform> params: Params;
-#endif
-#endif
-#endif
 
 #ifdef INPLACE
 fn inter_value(i: u32) -> f32 {
@@ -242,4 +208,3 @@ fn main(@builtin(workgroup_id) wid: vec3<u32>,
         col += WG_SIZE;
     }
 }
-
index 9d5d902cb1e240ed077ebfc9282478b722890f5f..c01df92f016478d225fdc31d8f0912c96ba361da 100644 (file)
@@ -29,7 +29,6 @@ struct Params {
 
     k: u32,
     ne2: u32,
-    ne3: u32,
 };
 
 @group(0) @binding(3)
index 66bfdd64015c0706bef81159fae5d483594c4f89..2d4c4e5a0b9186dfc8e0a06b46a19637efabb6b7 100644 (file)
@@ -39,7 +39,6 @@ struct Params {
     n_head: u32,
     n_group: u32,
     n_seq_tokens: u32,
-    n_seqs: u32,
 
     y_elems: u32,
 };