diff --git a/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp b/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp index 99d775c576..babaddb654 100644 --- a/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp +++ b/ggml/src/ggml-webgpu/ggml-webgpu-shader-lib.hpp @@ -591,7 +591,8 @@ struct ggml_webgpu_flash_attn_common_pipeline_key { ggml_type dst_type; uint32_t head_dim_qk; uint32_t head_dim_v; - bool kv_direct; + bool k_direct; + bool v_direct; bool kv_overlap; bool has_mask; bool has_sinks; @@ -600,8 +601,9 @@ struct ggml_webgpu_flash_attn_common_pipeline_key { bool operator==(const ggml_webgpu_flash_attn_common_pipeline_key & other) const { return q_type == other.q_type && k_type == other.k_type && v_type == other.v_type && dst_type == other.dst_type && head_dim_qk == other.head_dim_qk && head_dim_v == other.head_dim_v && - kv_direct == other.kv_direct && kv_overlap == other.kv_overlap && has_mask == other.has_mask && - has_sinks == other.has_sinks && uses_logit_softcap == other.uses_logit_softcap; + k_direct == other.k_direct && v_direct == other.v_direct && kv_overlap == other.kv_overlap && + has_mask == other.has_mask && has_sinks == other.has_sinks && + uses_logit_softcap == other.uses_logit_softcap; } }; @@ -613,7 +615,8 @@ inline void ggml_webgpu_flash_attn_hash_common_pipeline_key(size_t & ggml_webgpu_hash_combine(seed, key.dst_type); ggml_webgpu_hash_combine(seed, key.head_dim_qk); ggml_webgpu_hash_combine(seed, key.head_dim_v); - ggml_webgpu_hash_combine(seed, key.kv_direct); + ggml_webgpu_hash_combine(seed, key.k_direct); + ggml_webgpu_hash_combine(seed, key.v_direct); ggml_webgpu_hash_combine(seed, key.kv_overlap); ggml_webgpu_hash_combine(seed, key.has_mask); ggml_webgpu_hash_combine(seed, key.has_sinks); @@ -687,12 +690,13 @@ inline bool ggml_webgpu_flash_attn_float_vec4_aligned(const ggml_tensor * K, ggml_webgpu_flash_attn_float_vec4_aligned(V, storage_offset_alignment); } -inline bool ggml_webgpu_flash_attn_kv_direct(const ggml_tensor * Q, - const ggml_tensor * K, - const ggml_tensor * V, - uint32_t kv_direct_align) { - return K->type == GGML_TYPE_F16 && V->type == GGML_TYPE_F16 && (Q->ne[0] % kv_direct_align == 0) && - (K->ne[1] % GGML_WEBGPU_KV_SEQ_PAD == 0); +inline bool ggml_webgpu_flash_attn_k_direct(const ggml_tensor * Q, const ggml_tensor * K, uint32_t kv_direct_align) { + return (K->type == GGML_TYPE_F16 || K->type == GGML_TYPE_Q8_0 || K->type == GGML_TYPE_Q4_0) && + (Q->ne[0] % kv_direct_align == 0) && (K->ne[1] % GGML_WEBGPU_KV_SEQ_PAD == 0); +} + +inline bool ggml_webgpu_flash_attn_v_direct(const ggml_tensor * Q, const ggml_tensor * V, uint32_t kv_direct_align) { + return ggml_webgpu_flash_attn_k_direct(Q, V, kv_direct_align); } inline ggml_webgpu_flash_attn_common_pipeline_key ggml_webgpu_flash_attn_make_common_pipeline_key( @@ -706,10 +710,11 @@ inline ggml_webgpu_flash_attn_common_pipeline_key ggml_webgpu_flash_attn_make_co key.dst_type = context.dst->type; key.head_dim_qk = (uint32_t) context.src0->ne[0]; key.head_dim_v = (uint32_t) context.src2->ne[0]; - key.kv_direct = ggml_webgpu_flash_attn_kv_direct(context.src0, context.src1, context.src2, kv_direct_align); - key.kv_overlap = kv_overlap; - key.has_mask = context.src3 != nullptr; - key.has_sinks = context.src4 != nullptr; + key.k_direct = ggml_webgpu_flash_attn_k_direct(context.src0, context.src1, kv_direct_align); + key.v_direct = ggml_webgpu_flash_attn_v_direct(context.src0, context.src2, kv_direct_align); + key.kv_overlap = kv_overlap; + key.has_mask = context.src3 != nullptr; + key.has_sinks = context.src4 != nullptr; key.uses_logit_softcap = ggml_get_op_params_f32(context.dst, 2) != 0.0f; return key; } @@ -794,9 +799,13 @@ inline std::vector ggml_webgpu_flash_attn_common_defines( defines.push_back("LOGIT_SOFTCAP"); variant += "_lgsc"; } - if (key.kv_direct) { - defines.push_back("KV_DIRECT"); - variant += "_kvdirect"; + if (key.k_direct) { + defines.push_back("K_DIRECT"); + variant += "_k_direct"; + } + if (key.v_direct) { + defines.push_back("V_DIRECT"); + variant += "_v_direct"; } if (key.kv_overlap) { defines.push_back("KV_OVERLAP"); @@ -815,6 +824,12 @@ inline std::vector ggml_webgpu_flash_attn_common_defines( if (ggml_is_quantized(key.k_type) || ggml_is_quantized(key.v_type)) { defines.push_back("U32_DEQUANT_HELPERS"); + if (ggml_is_quantized(key.k_type)) { + defines.push_back("LOADERS_QUANTIZED_K"); + } + if (ggml_is_quantized(key.v_type)) { + defines.push_back("LOADERS_QUANTIZED_V"); + } } return defines; @@ -2792,12 +2807,14 @@ class ggml_webgpu_shader_lib { ggml_webgpu_flash_attn_pipeline_key key = {}; key.common = ggml_webgpu_flash_attn_make_common_pipeline_key( context, decisions.use_sg_matrix ? context.sg_mat_k : 1u, kv_overlap); - key.common.kv_direct = decisions.use_sg_matrix && key.common.kv_direct; - key.use_sg_matrix = decisions.use_sg_matrix; + key.common.k_direct &= decisions.use_sg_matrix && key.common.k_type == GGML_TYPE_F16; + key.common.v_direct &= decisions.use_sg_matrix && key.common.v_type == GGML_TYPE_F16; + key.use_sg_matrix = decisions.use_sg_matrix; const uint32_t max_kv_tile = ggml_webgpu_flash_attn_max_kv_tile( context.wg_mem_limit_bytes, decisions.q_tile, decisions.use_sg_matrix ? context.sg_mat_n : 1u, - key.common.head_dim_qk, key.common.head_dim_v, key.common.has_mask, key.common.kv_direct); + key.common.head_dim_qk, key.common.head_dim_v, key.common.has_mask, + key.common.k_direct || key.common.v_direct); GGML_ASSERT(max_kv_tile > 0); decisions.kv_tile = decisions.use_sg_matrix ? @@ -2809,7 +2826,7 @@ class ggml_webgpu_shader_lib { std::min(context.max_wg_size, std::max(GGML_WEBGPU_FLASH_ATTN_PREFERRED_WG_SIZE, GGML_WEBGPU_FLASH_ATTN_TILE_Q_TILE * context.max_subgroup_size)); - if (key.common.kv_direct) { + if (key.common.k_direct || key.common.v_direct) { decisions.kv_tile = std::min(decisions.kv_tile, GGML_WEBGPU_KV_SEQ_PAD); while (GGML_WEBGPU_KV_SEQ_PAD % decisions.kv_tile != 0) { decisions.kv_tile -= decisions.use_sg_matrix ? context.sg_mat_n : context.min_subgroup_size; @@ -2856,9 +2873,9 @@ class ggml_webgpu_shader_lib { } ggml_webgpu_flash_attn_vec_decisions decisions = {}; - decisions.kv_tile = - ggml_webgpu_flash_attn_get_vec_kv_tile(context.wg_mem_limit_bytes, key.common.head_dim_qk, - key.common.head_dim_v, key.common.has_mask, key.common.kv_direct); + decisions.kv_tile = ggml_webgpu_flash_attn_get_vec_kv_tile(context.wg_mem_limit_bytes, key.common.head_dim_qk, + key.common.head_dim_v, key.common.has_mask, + key.common.k_direct || key.common.v_direct); decisions.wg_size = context.max_subgroup_size; std::string variant = "flash_attn_vec"; @@ -2870,12 +2887,10 @@ class ggml_webgpu_shader_lib { variant += "_mask_blk"; } - uint32_t d_split = context.min_subgroup_size; - if (key.common.k_type == GGML_TYPE_F16 && key.common.v_type == GGML_TYPE_F16) { - const uint32_t D = key.common.head_dim_qk | key.common.head_dim_v; - const uint32_t D_lsb = D & (~(D - 1u)); - d_split = std::min(std::min(context.min_subgroup_size, 4u), std::max(D_lsb / 4u, 1u)); - } + uint32_t d_split = context.min_subgroup_size; + const uint32_t D = key.common.head_dim_qk | key.common.head_dim_v; + const uint32_t D_lsb = D & (~(D - 1u)); + d_split = std::min(std::min(context.min_subgroup_size, 4u), std::max(D_lsb / 4u, 1u)); defines.push_back(std::string("D_SPLIT=") + std::to_string(d_split)); variant += "_dsplit" + std::to_string(d_split); diff --git a/ggml/src/ggml-webgpu/ggml-webgpu.cpp b/ggml/src/ggml-webgpu/ggml-webgpu.cpp index 2add5da0b4..370f05dfe6 100644 --- a/ggml/src/ggml-webgpu/ggml-webgpu.cpp +++ b/ggml/src/ggml-webgpu/ggml-webgpu.cpp @@ -3839,7 +3839,8 @@ static size_t ggml_backend_webgpu_buffer_type_get_alloc_size(ggml_backend_buffer const auto & capabilities = ctx->webgpu_global_ctx->capabilities; if (ggml_webgpu_flash_attn_use_vec_path(ctx->webgpu_global_ctx, Q, K, V)) { const bool kv_direct = - ggml_webgpu_flash_attn_kv_direct(Q, K, V, GGML_WEBGPU_FLASH_ATTN_TILE_KV_VEC_WIDTH); + ggml_webgpu_flash_attn_k_direct(Q, K, GGML_WEBGPU_FLASH_ATTN_TILE_KV_VEC_WIDTH) || + ggml_webgpu_flash_attn_v_direct(Q, V, GGML_WEBGPU_FLASH_ATTN_TILE_KV_VEC_WIDTH); const uint32_t kv_tile = ggml_webgpu_flash_attn_get_vec_kv_tile( capabilities.limits.maxComputeWorkgroupStorageSize, (uint32_t) Q->ne[0], (uint32_t) V->ne[0], mask != nullptr, kv_direct); @@ -4448,9 +4449,10 @@ static bool ggml_backend_webgpu_device_supports_op(ggml_backend_dev_t dev, const const uint32_t q_tile = use_subgroup_matrix ? capabilities.sg_mat_m : GGML_WEBGPU_FLASH_ATTN_TILE_Q_TILE; const uint32_t kv_granularity = use_subgroup_matrix ? capabilities.sg_mat_n : 1u; - const bool kv_direct = use_subgroup_matrix ? - ggml_webgpu_flash_attn_kv_direct(src0, src1, src2, capabilities.sg_mat_k) : - false; + const bool kv_direct = use_subgroup_matrix ? + ggml_webgpu_flash_attn_k_direct(src0, src1, capabilities.sg_mat_k) || + ggml_webgpu_flash_attn_v_direct(src0, src2, capabilities.sg_mat_k) : + false; const uint32_t max_kv_tile = ggml_webgpu_flash_attn_max_kv_tile( capabilities.limits.maxComputeWorkgroupStorageSize, q_tile, kv_granularity, (uint32_t) src0->ne[0], (uint32_t) src2->ne[0], op->src[3] != nullptr, kv_direct); diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl b/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl index 6634fbd657..b0cf2853e0 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/common_decls.tmpl @@ -9,6 +9,12 @@ fn get_byte_i32(value: u32, index: u32) -> i32 { #endif #ifdef U32_DEQUANT_HELPERS + +fn f16_from_u16(bits: u32) -> f16 { + let packed = unpack2x16float(bits); + return f16(packed[0]); +} + #ifdef DECLARE_BYTE_LOADERS_SRC fn load_u16_at_src(byte_offset: u32) -> u32 { let word = src[byte_offset / 4u]; @@ -36,7 +42,7 @@ fn load_f16_as_f32_at_src(byte_offset: u32) -> f32 { let d_bits = (word >> shift) & 0xFFFFu; return unpack2x16float(d_bits)[0]; } -#endif +#endif // DECLARE_BYTE_LOADERS_SRC #ifdef DECLARE_BYTE_LOADERS_SRC0 fn load_u16_at_src0(byte_offset: u32) -> u32 { @@ -72,8 +78,47 @@ fn load_f16_as_f32_at_src0(byte_offset: u32) -> f32 { let d_bits = (word >> shift) & 0xFFFFu; return unpack2x16float(d_bits)[0]; } -#endif -#endif +#endif // DECLARE_BYTE_LOADERS_SRC0 + +#ifdef LOADERS_QUANTIZED_K +fn load_k_u16_at(byte_offset: u32) -> u32 { + let word = K[byte_offset / 4u]; + let shift = (byte_offset & 2u) * 8u; + return (word >> shift) & 0xFFFFu; +} + +fn load_k_u32_at(byte_offset: u32) -> u32 { + let word_idx = byte_offset / 4u; + let shift = (byte_offset & 3u) * 8u; + let lo = K[word_idx]; + if (shift == 0u) { + return lo; + } + let hi = K[word_idx + 1u]; + return (lo >> shift) | (hi << (32u - shift)); +} +#endif // LOADERS_QUANTIZED_K + +#ifdef LOADERS_QUANTIZED_V +fn load_v_u16_at(byte_offset: u32) -> u32 { + let word = V[byte_offset / 4u]; + let shift = (byte_offset & 2u) * 8u; + return (word >> shift) & 0xFFFFu; +} + +fn load_v_u32_at(byte_offset: u32) -> u32 { + let word_idx = byte_offset / 4u; + let shift = (byte_offset & 3u) * 8u; + let lo = V[word_idx]; + if (shift == 0u) { + return lo; + } + let hi = V[word_idx + 1u]; + return (lo >> shift) | (hi << (32u - shift)); +} +#endif // LOADERS_QUANTIZED_V + +#endif // U32_DEQUANT_HELPERS diff --git a/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn.wgsl b/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn.wgsl index 9767ca3d75..75f33e68ae 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn.wgsl @@ -138,7 +138,7 @@ const FLOAT_MIN: f32 = -1.0e9; // The number of Q rows processed per workgroup var q_shmem: array; -#ifndef KV_DIRECT +#if !defined(K_DIRECT) || !defined(V_DIRECT) 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; @@ -183,13 +183,12 @@ fn load_kx4(buf: ptr>, read_write>, scalar_index: u3 return (*buf)[scalar_index >> 2u]; } -#ifndef KV_DIRECT +#if !defined(K_DIRECT) || !defined(V_DIRECT) #define QUANT_SHMEM kv_shmem #define QUANT_OUT_TYPE f16 -#include "quant_inner_loops.tmpl" #include "flash_attn_quant_staging.tmpl" -#if !defined(K_Q4_0) && !defined(K_Q8_0) +#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; @@ -204,7 +203,7 @@ fn load_k_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, k_head_offset: u } #endif -#if !defined(V_Q4_0) && !defined(V_Q8_0) +#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; @@ -296,7 +295,7 @@ fn main(@builtin(workgroup_id) wg_id: vec3, } // load k tile into shared memory -#ifndef KV_DIRECT +#ifndef K_DIRECT load_k_tile_block(local_id.x, kv_count, kv_tile, k_head_offset); #endif @@ -306,7 +305,7 @@ fn main(@builtin(workgroup_id) wg_id: vec3, // TODO: this loop seems to be the current largest bottleneck // this bracket exists to scope the lifetime of variables, reducing register pressure { -#ifdef KV_DIRECT +#ifdef K_DIRECT let k_block_row = kv_tile + subgroup_id * SG_MAT_N; var k_global_offset = k_head_offset + k_block_row * params.stride_k1; #else @@ -318,7 +317,7 @@ fn main(@builtin(workgroup_id) wg_id: vec3, var q_cur = subgroupMatrixLoad>(&q_shmem, 0u, false, HEAD_DIM_QK); -#ifdef KV_DIRECT +#ifdef K_DIRECT var k_cur = subgroupMatrixLoad>(&K, k_global_offset + 0u, true, params.stride_k1); #else var k_cur = subgroupMatrixLoad>(&kv_shmem, k_block_offset + 0u, true, HEAD_DIM_QK); @@ -328,7 +327,7 @@ fn main(@builtin(workgroup_id) wg_id: vec3, for (; t + 1u < HEAD_DIM_QK / SG_MAT_K; t += 2u) { let h0 = t * SG_MAT_K; var q0 = subgroupMatrixLoad>(&q_shmem, h0, false, HEAD_DIM_QK); -#ifdef KV_DIRECT +#ifdef K_DIRECT var k0 = subgroupMatrixLoad>(&K, k_global_offset + h0, true, params.stride_k1); #else var k0 = subgroupMatrixLoad>(&kv_shmem, k_block_offset + h0, true, HEAD_DIM_QK); @@ -339,7 +338,7 @@ fn main(@builtin(workgroup_id) wg_id: vec3, let h1 = (t + 1u) * SG_MAT_K; var q1g = subgroupMatrixLoad>(&q_shmem, h1, false, HEAD_DIM_QK); -#ifdef KV_DIRECT +#ifdef K_DIRECT var k1g = subgroupMatrixLoad>(&K, k_global_offset + h1, true, params.stride_k1); #else var k1g = subgroupMatrixLoad>(&kv_shmem, k_block_offset + h1, true, HEAD_DIM_QK); @@ -353,7 +352,7 @@ fn main(@builtin(workgroup_id) wg_id: vec3, if (t < HEAD_DIM_QK / SG_MAT_K) { let h = t * SG_MAT_K; var qn = subgroupMatrixLoad>(&q_shmem, h, false, HEAD_DIM_QK); -#ifdef KV_DIRECT +#ifdef K_DIRECT var kn = subgroupMatrixLoad>(&K, k_global_offset + h, true, params.stride_k1); #else var kn = subgroupMatrixLoad>(&kv_shmem, k_block_offset + h, true, HEAD_DIM_QK); @@ -365,7 +364,7 @@ fn main(@builtin(workgroup_id) wg_id: vec3, acc = subgroupMatrixMultiplyAccumulate(q_cur, k_cur, acc); -#ifdef KV_DIRECT +#ifdef K_DIRECT k_global_offset += num_subgroups * SG_MAT_N * params.stride_k1; #else k_block_offset += num_subgroups * SG_MAT_N * HEAD_DIM_QK; @@ -436,7 +435,7 @@ fn main(@builtin(workgroup_id) wg_id: vec3, } // load v tile into shared memory -#ifndef KV_DIRECT +#ifndef V_DIRECT load_v_tile_block(local_id.x, kv_count, kv_tile, v_head_offset); #endif @@ -464,7 +463,7 @@ fn main(@builtin(workgroup_id) wg_id: vec3, ); // load V submatrix from global or shared memory -#ifdef KV_DIRECT +#ifdef V_DIRECT let v_block_row = kv_tile + kv_block * SG_MAT_N; let v_global_offset = v_head_offset + v_block_row * params.stride_v1 + head_dim_block; var v_sg_mat: subgroup_matrix_right = subgroupMatrixLoad>( 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 index 8f41eb7bfd..1c23260df0 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_quant_staging.tmpl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_quant_staging.tmpl @@ -1,3 +1,5 @@ +#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) @@ -26,49 +28,6 @@ #define V_BYTES_PER_INNER_LOOP 4u #endif -#if defined(K_Q4_0) || defined(K_Q8_0) -fn load_k_u16_at(byte_offset: u32) -> u32 { - let word = K[byte_offset / 4u]; - let shift = (byte_offset & 2u) * 8u; - return (word >> shift) & 0xFFFFu; -} - -fn load_k_u32_at(byte_offset: u32) -> u32 { - let word_idx = byte_offset / 4u; - let shift = (byte_offset & 3u) * 8u; - let lo = K[word_idx]; - if (shift == 0u) { - return lo; - } - let hi = K[word_idx + 1u]; - return (lo >> shift) | (hi << (32u - shift)); -} -#endif - -#if defined(V_Q4_0) || defined(V_Q8_0) -fn load_v_u16_at(byte_offset: u32) -> u32 { - let word = V[byte_offset / 4u]; - let shift = (byte_offset & 2u) * 8u; - return (word >> shift) & 0xFFFFu; -} - -fn load_v_u32_at(byte_offset: u32) -> u32 { - let word_idx = byte_offset / 4u; - let shift = (byte_offset & 3u) * 8u; - let lo = V[word_idx]; - if (shift == 0u) { - return lo; - } - let hi = V[word_idx + 1u]; - return (lo >> shift) | (hi << (32u - shift)); -} -#endif - -fn f16_from_u16(bits: u32) -> f16 { - let packed = unpack2x16float(bits); - return f16(packed[0]); -} - #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) { 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 e68934113f..43f4fe7cac 100644 --- a/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_tile.wgsl +++ b/ggml/src/ggml-webgpu/wgsl-shaders/flash_attn_tile.wgsl @@ -153,7 +153,6 @@ var p_shmem: array; #define QUANT_SHMEM kv_shmem #define QUANT_OUT_TYPE f16 -#include "quant_inner_loops.tmpl" #include "flash_attn_quant_staging.tmpl" #if !defined(K_Q4_0) && !defined(K_Q8_0) @@ -270,7 +269,9 @@ fn main(@builtin(workgroup_id) wg_id: vec3, local_scores[slot] = FLOAT_MIN; } -#ifndef KV_DIRECT + // The tile path stages K/V in shared memory so each tile can be reused across + // Q_TILE query rows. It therefore does not use the direct path. +#ifndef K_DIRECT load_k_tile_block(local_id.x, kv_count, kv_tile, k_head_offset); #endif @@ -333,7 +334,9 @@ fn main(@builtin(workgroup_id) wg_id: vec3, workgroupBarrier(); -#ifndef KV_DIRECT + // The tile path stages K/V in shared memory so each tile can be reused across + // Q_TILE query rows. It therefore does not use the direct path. +#ifndef V_DIRECT load_v_tile_block(local_id.x, kv_count, kv_tile, v_head_offset); #endif 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 d512762419..b8e0be90d9 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 @@ -196,49 +196,35 @@ struct Params { // Just a very small float value. const FLOAT_MIN: f32 = -1.0e9; +const kv_shmem_size = KV_TILE * max(HEAD_DIM_QK, HEAD_DIM_V); var q_shmem: array; - -#ifndef KV_DIRECT -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; -#endif - var o_shmem: array; +// note that we reuse the same storage for both since we only need one at a time +var inter_shmem: array; #ifdef MASK // storage for mask values var mask_shmem: array; #endif -// note that we reuse the same storage for both since we only need one at a time -var inter_shmem: array; - -// 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, - inter_shmem[kv_idx] * params.scale, - kv_idx < KV_TILE); -#ifdef LOGIT_SOFTCAP - v = params.logit_softcap * tanh(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 -#ifdef MASK - if (apply_mask) { - var mask_val = select(0.0, mask_shmem[kv_idx], kv_idx < KV_TILE); - v += select(mask_val, slope * mask_val, has_bias); - } -#endif - return v; -} -#ifndef KV_DIRECT +// 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 "quant_inner_loops.tmpl" #include "flash_attn_quant_staging.tmpl" -#if !defined(K_Q4_0) && !defined(K_Q8_0) +#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; @@ -256,7 +242,7 @@ fn load_k_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, k_head_offset: u } #endif -#if !defined(V_Q4_0) && !defined(V_Q8_0) +#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; @@ -273,7 +259,24 @@ fn load_v_tile_block(local_x: u32, kv_count: u32, kv_tile: u32, v_head_offset: u } } #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, + inter_shmem[kv_idx] * params.scale, + kv_idx < KV_TILE); +#ifdef LOGIT_SOFTCAP + v = params.logit_softcap * tanh(v); #endif +#ifdef MASK + if (apply_mask) { + var mask_val = select(0.0, mask_shmem[kv_idx], kv_idx < KV_TILE); + v += select(mask_val, slope * mask_val, has_bias); + } +#endif + return v; +} @compute @workgroup_size(WG_SIZE) fn main(@builtin(workgroup_id) wg_id: vec3, @@ -355,12 +358,31 @@ fn main(@builtin(workgroup_id) wg_id: vec3, inter_shmem[elem_idx] = 0.0; } - // load k tile into shared memory -#ifndef KV_DIRECT - load_k_tile_block(local_id.x, kv_count, kv_tile, k_head_offset); +#ifdef K_DIRECT + // load only the scale factor (d) from each quantized block into shared memory on the direct path. +#if defined(K_Q8_0) + for (var j = local_id.x * 32; j < KV_TILE * HEAD_DIM_QK; j += WG_SIZE * 32) { + let kv_row = kv_tile + j / HEAD_DIM_QK; + let block_idx = (j % HEAD_DIM_QK) / 32; + let block_byte_base = 34 * (k_head_offset + kv_row * params.stride_k1 + block_idx); + let d = f32(f16_from_u16(load_k_u16_at(block_byte_base))); + d_shmem[j / 32] = d; + } +#elif defined(K_Q4_0) + for (var j = local_id.x * 32; j < KV_TILE * HEAD_DIM_QK; j += WG_SIZE * 32) { + let kv_row = kv_tile + j / HEAD_DIM_QK; + let block_idx = (j % HEAD_DIM_QK) / 32; + let block_byte_base = 18 * (k_head_offset + kv_row * params.stride_k1 + block_idx); + let d = f32(f16_from_u16(load_k_u16_at(block_byte_base))); + d_shmem[j / 32] = d; + } #endif +#else + // load k tile into shared memory + load_k_tile_block(local_id.x, kv_count, kv_tile, k_head_offset); +#endif // defined(K_DIRECT) - workgroupBarrier(); + workgroupBarrier(); // accumulate q block * k block into registers across the entire KV tile if (!skip_tile) { @@ -381,9 +403,40 @@ fn main(@builtin(workgroup_id) wg_id: vec3, q_shmem[q_off + 1u], q_shmem[q_off + 2u], q_shmem[q_off + 3u]); -#ifdef KV_DIRECT +#ifdef K_DIRECT +#if defined(K_Q8_0) + let kv_row = kv_tile + kv_idx; + let block_idx = (i * 4u) / 32; + let id_in_block = (i * 4u) % 32; + let block_byte_base = 34 * (k_head_offset + kv_row * params.stride_k1 + block_idx); + let q_byte_base = block_byte_base + 2u; + let d = d_shmem[(kv_idx * HEAD_DIM_QK) / 32 + block_idx]; + let q8u4 = load_k_u32_at(q_byte_base + id_in_block); + let kv = vec4( + d * f32(get_byte_i32(q8u4, 0)), + d * f32(get_byte_i32(q8u4, 1)), + d * f32(get_byte_i32(q8u4, 2)), + d * f32(get_byte_i32(q8u4, 3)), + ); +#elif defined(K_Q4_0) + let kv_row = kv_tile + kv_idx; + let block_idx = (i * 4u) / 32; + let id_in_block = (i * 4u) % 32; + let phase = id_in_block / 16; + let block_byte_base = 18 * (k_head_offset + kv_row * params.stride_k1 + block_idx); + let q_byte_base = block_byte_base + 2u; + let d = d_shmem[(kv_idx * HEAD_DIM_QK) / 32 + block_idx]; + let q8u4 = load_k_u32_at(q_byte_base + (id_in_block - phase * 16u)); + let kv = vec4( + d * (f32((get_byte(q8u4, 0) >> (phase * 4u)) & 0xFu) - 8.0), + d * (f32((get_byte(q8u4, 1) >> (phase * 4u)) & 0xFu) - 8.0), + d * (f32((get_byte(q8u4, 2) >> (phase * 4u)) & 0xFu) - 8.0), + d * (f32((get_byte(q8u4, 3) >> (phase * 4u)) & 0xFu) - 8.0), + ); +#else let idx = k_head_offset + (kv_tile + kv_idx) * params.stride_k1 + (i * 4u); let kv = vec4(K[idx >> 2u]); +#endif #else let idx = kv_idx * HEAD_DIM_QK + (i * 4u); let kv = vec4( @@ -391,7 +444,7 @@ fn main(@builtin(workgroup_id) wg_id: vec3, kv_shmem[idx + 1u], kv_shmem[idx + 2u], kv_shmem[idx + 3u]); -#endif +#endif // defined(K_DIRECT) partial_sum += dot(qv, kv); } } @@ -473,12 +526,32 @@ fn main(@builtin(workgroup_id) wg_id: vec3, } } - // load v tile into shared memory -#ifndef KV_DIRECT - load_v_tile_block(local_id.x, kv_count, kv_tile, v_head_offset); -#endif - workgroupBarrier(); +#ifdef V_DIRECT + // load only `d` of quantized block into shared memory in the direct path +#if defined(V_Q8_0) + for (var j = local_id.x * 32; j < KV_TILE * HEAD_DIM_V; j += WG_SIZE * 32) { + let v_row = kv_tile + j / HEAD_DIM_V; + let block_idx = (j % HEAD_DIM_V) / 32; + let block_byte_base = 34 * (v_head_offset + v_row * params.stride_v1 + block_idx); + let d = f32(f16_from_u16(load_v_u16_at(block_byte_base))); + d_shmem[j / 32] = d; + } +#elif defined(V_Q4_0) + for (var j = local_id.x * 32; j < KV_TILE * HEAD_DIM_V; j += WG_SIZE * 32) { + let v_row = kv_tile + j / HEAD_DIM_V; + let block_idx = (j % HEAD_DIM_V) / 32; + let block_byte_base = 18 * (v_head_offset + v_row * params.stride_v1 + block_idx); + let d = f32(f16_from_u16(load_v_u16_at(block_byte_base))); + d_shmem[j / 32] = d; + } +#endif +#else + // load v tile into shared memory + load_v_tile_block(local_id.x, kv_count, kv_tile, v_head_offset); +#endif // V_DIRECT + + workgroupBarrier(); if (!skip_tile) { // we have P (KV_TILE) in inter_shmem and V (KV_TILE x head_dim_v) in kv_shmem @@ -501,9 +574,38 @@ fn main(@builtin(workgroup_id) wg_id: vec3, } let p = inter_shmem[kv_idx]; -#ifdef KV_DIRECT +#ifdef V_DIRECT +#if defined(V_Q8_0) + let block_idx = (vec_col * 4u) / 32; + let id_in_block = (vec_col * 4u) % 32; + let block_byte_base = 34 * (v_head_offset + v_row * params.stride_v1 + block_idx); + let q_byte_base = block_byte_base + 2u; + let d = d_shmem[(kv_idx * HEAD_DIM_V) / 32 + block_idx]; + let q8u4 = load_v_u32_at(q_byte_base + id_in_block); + let v4 = vec4( + d * f32(get_byte_i32(q8u4, 0)), + d * f32(get_byte_i32(q8u4, 1)), + d * f32(get_byte_i32(q8u4, 2)), + d * f32(get_byte_i32(q8u4, 3)), + ); +#elif defined(V_Q4_0) + let block_idx = (vec_col * 4u) / 32; + let id_in_block = (vec_col * 4u) % 32; + let phase = id_in_block / 16; + let block_byte_base = 18 * (v_head_offset + v_row * params.stride_v1 + block_idx); + let q_byte_base = block_byte_base + 2u; + let d = d_shmem[(kv_idx * HEAD_DIM_V) / 32 + block_idx]; + let q8u4 = load_v_u32_at(q_byte_base + (id_in_block - phase * 16u)); + let v4 = vec4( + d * (f32((get_byte(q8u4, 0) >> (phase * 4u)) & 0xFu) - 8.0), + d * (f32((get_byte(q8u4, 1) >> (phase * 4u)) & 0xFu) - 8.0), + d * (f32((get_byte(q8u4, 2) >> (phase * 4u)) & 0xFu) - 8.0), + d * (f32((get_byte(q8u4, 3) >> (phase * 4u)) & 0xFu) - 8.0), + ); +#else let v_idx = v_head_offset + v_row * params.stride_v1 + vec_col * 4u; let v4 = vec4(V[v_idx >> 2u]); +#endif #else let v_idx = kv_idx * HEAD_DIM_V + vec_col * 4u; let v4 = vec4( @@ -511,7 +613,7 @@ fn main(@builtin(workgroup_id) wg_id: vec3, kv_shmem[v_idx + 1u], kv_shmem[v_idx + 2u], kv_shmem[v_idx + 3u]); -#endif +#endif // defined(V_DIRECT) lo += p * v4; }