From 20a6246f290b1d4dd21abd10805fa3c585ad36a5 Mon Sep 17 00:00:00 2001 From: Jeff Bolz Date: Thu, 8 May 2025 14:55:52 -0500 Subject: [PATCH] vulkan: avoid using Float16 capability in scalar FA --- ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp | 6 +++--- 1 file changed, 3 insertions(+), 3 deletions(-) diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp index a055e9929..e6545160d 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp @@ -137,7 +137,7 @@ ACC_TYPE perElemOpComputeSlope(const in uint32_t r, const in uint32_t c, const i shared FLOAT_TYPE tmpsh[gl_WorkGroupSize.x]; shared vec4 tmpshv4[gl_WorkGroupSize.x]; -shared float16_t masksh[Bc][Br]; +shared float masksh[Bc][Br]; shared vec4 Qf[Br][D / 4]; void main() { @@ -296,14 +296,14 @@ void main() { uint32_t c = (idx + tid) % Bc; uint32_t r = (idx + tid) / Bc; if (idx + tid < Bc * Br) { - masksh[c][r] = data_m[(i * Br + r) * m_stride + (j * Bc + c)]; + masksh[c][r] = float(data_m[(i * Br + r) * m_stride + (j * Bc + c)]); } } barrier(); [[unroll]] for (uint32_t c = 0; c < cols_per_thread; ++c) { [[unroll]] for (uint32_t r = 0; r < Br; ++r) { - float mvf = float(masksh[c * cols_per_iter + col_tid][r]); + float mvf = masksh[c * cols_per_iter + col_tid][r]; Sf[r][c] += slope[r]*mvf; }