Commit 60081bb2b for llama.cpp

commit 60081bb2b5b3294165a4d67c5cbeebe74c868014
Author: dsproule <dsproule@qti.qualcomm.com>
Date:   Fri Sep 18 16:32:31 2026 -0700

    opencl: add support for bin kernel `flash_attn_f32_f16_bin` (#29046)

    * opencl: add `flash_attn_f32_f16_bin`

    * opencl: guarded prefill fa

diff --git a/ggml/src/ggml-opencl/CMakeLists.txt b/ggml/src/ggml-opencl/CMakeLists.txt
index 53e938618..ff5e8ef46 100644
--- a/ggml/src/ggml-opencl/CMakeLists.txt
+++ b/ggml/src/ggml-opencl/CMakeLists.txt
@@ -233,6 +233,7 @@ set(GGML_OPENCL_KERNELS
     mul_mm_f16_f32_kq_kqv
     conv2d
     conv2d_f16_f32
+    flash_attn_repack
     flash_attn_pre_f16
     flash_attn_f32_f16
     flash_attn_f32_q8_0
diff --git a/ggml/src/ggml-opencl/ggml-opencl.cpp b/ggml/src/ggml-opencl/ggml-opencl.cpp
index 1c26797b9..fe7377b2d 100644
--- a/ggml/src/ggml-opencl/ggml-opencl.cpp
+++ b/ggml/src/ggml-opencl/ggml-opencl.cpp
@@ -567,6 +567,16 @@ struct ggml_opencl_fa_kernels {
     // attempted (variant, (dk, dv))
     // all attempted FA kernels appear here, but those not registered failed compilation
     std::set<std::pair<int, std::pair<int, int>>> variant_attempted;
+
+    // FA bin kernels
+#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
+    cl_kernel kernel_flash_attn_f32_f16_bin;
+
+    cl_kernel kernel_repack_q_for_wmm;
+    cl_kernel kernel_repack_k_for_wmm;
+    cl_kernel kernel_repack_v_for_wmm;
+    cl_kernel kernel_repack_mask_for_wmm;
+#endif
 };

 #ifdef GGML_OPENCL_USE_ADRENO_KERNELS
@@ -5172,6 +5182,43 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) {
         CL_CHECK(clReleaseProgram(prog));
         GGML_LOG_CONT(".");
     }
+
+    // repack
+    {
+#ifdef GGML_OPENCL_EMBED_KERNELS
+        const std::string kernel_src {
+            #include "flash_attn_repack.cl.h"
+        };
+#else
+        const std::string kernel_src = read_file("flash_attn_repack.cl");
+#endif
+        cl_program prog =
+            build_program_from_source(backend_ctx, kernel_src.c_str(), compile_opts);
+
+        CL_CHECK((backend_ctx->fa.kernel_repack_q_for_wmm = clCreateKernel(prog, "kernel_repack_q_for_wmm", &err), err));
+        CL_CHECK((backend_ctx->fa.kernel_repack_k_for_wmm = clCreateKernel(prog, "kernel_repack_k_for_wmm", &err), err));
+        CL_CHECK((backend_ctx->fa.kernel_repack_v_for_wmm = clCreateKernel(prog, "kernel_repack_v_for_wmm", &err), err));
+        CL_CHECK((backend_ctx->fa.kernel_repack_mask_for_wmm = clCreateKernel(prog, "kernel_repack_mask_for_wmm", &err), err));
+        GGML_LOG_CONT(".");
+    }
+
+    // kernel_flash_attn_f32_f16_bin
+    {
+        size_t bin_size = 0;
+        backend_ctx->fa.kernel_flash_attn_f32_f16_bin = nullptr;
+
+        if (use_adreno_bin_kernels(backend_ctx)) {
+            const char * kernel_bin = (const char *)backend_ctx->get_adreno_bin_kernel("flash_attn_f32_f16_wmm", &bin_size);
+            if (kernel_bin && bin_size > 0) {
+                cl_program prog =
+                    build_program_from_binary(backend_ctx->context, backend_ctx->device, kernel_bin, CL_moe_compile_opts, bin_size);
+
+                CL_CHECK((backend_ctx->fa.kernel_flash_attn_f32_f16_bin = clCreateKernel(prog, "flash_attn_f32_f16", &err), err));
+                CL_CHECK(clReleaseProgram(prog));
+                GGML_LOG_CONT(".");
+            }
+        }
+    }
 #endif // GGML_OPENCL_USE_ADRENO_KERNELS
     GGML_LOG_CONT("\n");
     backend_ctx->kernels_loaded = true;
@@ -8532,6 +8579,28 @@ inline bool use_q4_0_bin_kernels(const ggml_backend_opencl_context *backend_ctx,
 #endif
 }

+#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
+static bool use_fa_bin_kernels_prefill(const ggml_backend_opencl_context * backend_ctx, const ggml_tensor * q, const ggml_tensor * k, const ggml_tensor * v) {
+    if (backend_ctx->fa.kernel_flash_attn_f32_f16_bin == nullptr) {
+        return false;
+    }
+
+    const bool is_mixed = q->type == GGML_TYPE_F32 && k->type == GGML_TYPE_F16 && v->type == GGML_TYPE_F16;
+    const bool is_q8_0 = q->type == GGML_TYPE_F32 && k->type == GGML_TYPE_Q8_0 && v->type == GGML_TYPE_Q8_0;
+
+    const int n_q = q->ne[1];
+    const int dk = q->ne[0];
+    const int dv = v->ne[0];
+
+    constexpr bool prefill_only = true;
+
+    return (backend_ctx->gpu_family == GPU_FAMILY::ADRENO &&
+            (is_mixed || is_q8_0) && (dk == dv)
+            && (dk == 64 || dk == 128 || dk == 256 || dk == 512)
+            && (!prefill_only || n_q != 1));
+}
+#endif
+
 // The flat-GEMV large-m escape is OPT-IN (GGML_OPENCL_FLAT_LARGE_M=1) because it
 // is SLOWER than the route it replaces, not because it is unsafe. It was first
 // parked on the theory that it out-of-bounds-writes at vocab-scale shapes; that
@@ -8990,6 +9059,11 @@ static bool ggml_opencl_supports_op(ggml_backend_dev_t dev, const struct ggml_te
         case GGML_OP_MEAN:
             return op->src[0]->type == GGML_TYPE_F32;
         case GGML_OP_FLASH_ATTN_EXT: {
+#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
+            if (use_fa_bin_kernels_prefill(backend_ctx, op->src[0], op->src[1], op->src[2])) {
+                return true;
+            }
+#endif
             // The E17 compilers segfault while building FA kernels, skip E17 for now
             if (adreno_e17_compiler_quirks(backend_ctx)) {
                 return false;
@@ -17198,6 +17272,407 @@ static void ggml_cl_adreno_xmem_attn_run(

 #endif // GGML_OPENCL_USE_ADRENO_KERNELS

+#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
+static void ggml_cl_flash_attn_prefill_bin(ggml_backend_t backend, const ggml_tensor * q, const ggml_tensor * k, ggml_tensor * dst) {
+    const ggml_tensor * v = dst->src[2];
+    const ggml_tensor * mask = dst->src[3];
+    const ggml_tensor * sinks = dst->src[4];
+    GGML_ASSERT(q->extra);
+    GGML_ASSERT(k->extra);
+    GGML_ASSERT(v->extra);
+    GGML_ASSERT(dst->extra);
+    if (mask) {
+        GGML_ASSERT(mask->extra);
+    }
+    if (sinks) {
+        GGML_ASSERT(sinks->extra);
+    }
+
+    ggml_backend_opencl_context *backend_ctx = (ggml_backend_opencl_context *)backend->context;
+    cl_context context = backend_ctx->context;
+
+    const int n_q = q->ne[1];
+    const int n_kv = k->ne[1];
+    const int d_head_q = q->ne[0];
+    const int d_head_v = v->ne[0];
+    const int n_head = q->ne[2];
+    const int n_head_kv = k->ne[2];
+    const int n_batch = q->ne[3];
+
+    const std::pair<int, int> dk_dv = {d_head_q, d_head_v};
+    cl_kernel kernel = backend_ctx->fa.kernel_flash_attn_f32_f16_bin;
+    GGML_ASSERT(kernel != NULL);
+
+    ggml_tensor_extra_cl * extra_q = (ggml_tensor_extra_cl *)q->extra;
+    ggml_tensor_extra_cl * extra_k = (ggml_tensor_extra_cl *)k->extra;
+    ggml_tensor_extra_cl * extra_v = (ggml_tensor_extra_cl *)v->extra;
+    ggml_tensor_extra_cl * extra_o = (ggml_tensor_extra_cl *)dst->extra;
+    ggml_tensor_extra_cl * extra_mask = mask ? (ggml_tensor_extra_cl *)mask->extra : NULL;
+    ggml_tensor_extra_cl * extra_sinks = sinks ? (ggml_tensor_extra_cl *)sinks->extra : NULL;
+
+    cl_ulong offset_q = extra_q->offset + q->view_offs;
+    cl_ulong offset_o = extra_o->offset + dst->view_offs;
+
+    cl_mem   mask_buffer = extra_mask ? extra_mask->data_device : NULL;
+    cl_ulong offset_mask = extra_mask ? extra_mask->offset + mask->view_offs : 0;
+    cl_mem   sinks_buffer = extra_sinks ? extra_sinks->data_device : NULL;
+    cl_ulong offset_sinks = extra_sinks ? extra_sinks->offset + sinks->view_offs : 0;
+
+    const cl_ulong q_nb1 = q->nb[1];
+    const cl_ulong q_nb2 = q->nb[2];
+    const cl_ulong q_nb3 = q->nb[3];
+
+    cl_mem   k_data_device = extra_k->data_device;
+    cl_ulong offset_k = extra_k->offset + k->view_offs;
+    cl_ulong k_nb1 = k->nb[1];
+    cl_ulong k_nb2 = k->nb[2];
+    cl_ulong k_nb3 = k->nb[3];
+
+    cl_mem   v_data_device = extra_v->data_device;
+    cl_ulong offset_v = extra_v->offset + v->view_offs;
+    cl_ulong v_nb1 = v->nb[1];
+    cl_ulong v_nb2 = v->nb[2];
+    cl_ulong v_nb3 = v->nb[3];
+
+    const cl_ulong o_nb1 = dst->nb[1];
+    const cl_ulong o_nb2 = dst->nb[2];
+    const cl_ulong o_nb3 = dst->nb[3];
+
+    const cl_ulong mask_nb1 = mask ? mask->nb[1] : 0;
+    const cl_ulong mask_nb2 = mask ? mask->nb[2] : 0;
+    const cl_ulong mask_nb3 = mask ? mask->nb[3] : 0;
+    const int mask_ne2 = mask ? mask->ne[2] : 0;
+    const int mask_ne3 = mask ? mask->ne[3] : 0;
+
+    float * params      = (float *)dst->op_params;
+    float scale         = params[0];
+    float max_bias      = params[1];
+    float logit_softcap = params[2];
+
+    const int is_causal = (mask == NULL && n_q > 1 && n_q == n_kv);     // redundant n_q > 1 check ?
+
+    const int n_head_log2_val = n_head > 0 ? 1u << (int)floorf(log2f((float)n_head)) : 0;
+    const float n_head_log2_f = n_head_log2_val > 0 ? (float)n_head_log2_val : 1.0f;
+    const float m0 = powf(2.0f, -(max_bias) / n_head_log2_f);
+    const float m1 = powf(2.0f, -(max_bias / 2.0f) / n_head_log2_f);
+
+    const bool is_q8_0 = q->type == GGML_TYPE_F32 && k->type == GGML_TYPE_Q8_0 && v->type == GGML_TYPE_Q8_0;
+
+    ggml_cl_flash_attn_temp_buffer temp_k;
+    ggml_cl_flash_attn_temp_buffer temp_v;
+    ggml_cl_flash_attn_temp_buffer temp_k_aos;
+    ggml_cl_flash_attn_temp_buffer temp_v_aos;
+
+    if (is_q8_0) {
+        ggml_cl_flash_attn_reconstruct_aos(
+            backend_ctx, k, temp_k_aos, k_data_device, offset_k, k_nb1, k_nb2, k_nb3);
+
+        ggml_cl_flash_attn_reconstruct_aos(
+            backend_ctx, v, temp_v_aos, v_data_device, offset_v, v_nb1, v_nb2, v_nb3);
+
+        bool k_done = ggml_cl_flash_attn_dequant_kv_gpu(
+            backend_ctx, k, GGML_TYPE_F16, k_data_device, offset_k, k_nb1, k_nb2, k_nb3,
+            temp_k, k_data_device, offset_k, k_nb1, k_nb2, k_nb3);
+
+        bool v_done = ggml_cl_flash_attn_dequant_kv_gpu(
+            backend_ctx, v, GGML_TYPE_F16, v_data_device, offset_v, v_nb1, v_nb2, v_nb3,
+            temp_v, v_data_device, offset_v, v_nb1, v_nb2, v_nb3);
+
+        GGML_ASSERT(k_done && v_done);
+    }
+
+    // Allocate input/output memory buffers
+    cl_mem mem_matrixQ;
+    cl_mem mem_matrixK;
+    cl_mem mem_matrixV;
+    cl_mem mem_matrixO;
+    cl_buffer_region region;
+    cl_int err;
+
+    region.origin = offset_q;
+    region.size = ggml_nbytes(q);
+    mem_matrixQ = clCreateSubBuffer(extra_q->data_device, CL_MEM_READ_WRITE, CL_BUFFER_CREATE_TYPE_REGION, &region, &err);
+    CL_CHECK(err);
+
+    region.origin = offset_k;
+    region.size = is_q8_0 ? (size_t) k_nb3 * (size_t) k->ne[3] : ggml_nbytes(k);
+    mem_matrixK = clCreateSubBuffer(k_data_device, CL_MEM_READ_WRITE, CL_BUFFER_CREATE_TYPE_REGION, &region, &err);
+    CL_CHECK(err);
+
+    region.origin = offset_v;
+    region.size = is_q8_0 ? (size_t) v_nb3 * (size_t) v->ne[3] : ggml_nbytes(v);
+    mem_matrixV = clCreateSubBuffer(v_data_device, CL_MEM_READ_WRITE, CL_BUFFER_CREATE_TYPE_REGION, &region, &err);
+    CL_CHECK(err);
+
+    region.origin = offset_o;
+    region.size = ggml_nbytes(dst);
+    mem_matrixO = clCreateSubBuffer(extra_o->data_device, CL_MEM_READ_WRITE, CL_BUFFER_CREATE_TYPE_REGION, &region, &err);
+    CL_CHECK(err);
+
+    cl_image_format img_fmt_1d = { CL_RGBA, CL_FLOAT};
+    cl_image_desc img_desc_1d;
+
+    // use image 1d buffer used as fallback when on mask is applied
+    cl_mem  mem_tex_mask_fallback_1dbuf;
+    img_fmt_1d = { CL_RGBA, CL_HALF_FLOAT};
+    memset(&img_desc_1d, 0, sizeof(img_desc_1d));
+    img_desc_1d.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER;
+    img_desc_1d.image_width = 1;
+    img_desc_1d.buffer = mem_matrixK;
+    mem_tex_mask_fallback_1dbuf = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt_1d, &img_desc_1d, NULL, &err);
+    CL_CHECK(err);
+
+    cl_mem  mem_tex_matrixO_1dbuf;
+    img_fmt_1d = { CL_RGBA, CL_FLOAT};
+    memset(&img_desc_1d, 0, sizeof(img_desc_1d));
+    img_desc_1d.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER;
+    img_desc_1d.image_width = ggml_nbytes(dst) / 4 / 4;
+    img_desc_1d.buffer = mem_matrixO;
+    mem_tex_matrixO_1dbuf = clCreateImage(context, CL_MEM_WRITE_ONLY, &img_fmt_1d, &img_desc_1d, NULL, &err);
+    CL_CHECK(err);
+
+    // The bin kernel requires 2d (or 3d) buffers packed for data loading/multiplication.
+    // These repack kernels launch across all buffers to ensure compatibility
+    cl_mem  mem_tex_matrixMask_1dbuf = NULL;
+    cl_mem  mem_matrixMask = NULL;
+    cl_mem  mem_matrixMask_padded = NULL;
+    cl_ulong mask_nb1_padded = mask_nb1, mask_nb2_padded = mask_nb2, mask_nb3_padded = mask_nb3;
+    if (extra_mask) {
+        // allocate mem_matrixMask w/ new padded size
+        size_t n_kv_padded = GGML_PAD(n_kv, 4);
+        size_t mask_nb_padded = n_kv_padded * sizeof(cl_half) * mask->ne[1] * mask->ne[2] * mask->ne[3];
+
+        // apply offset and create subBuffer for mask
+        region.origin = offset_mask;
+        region.size = ggml_nbytes(mask);
+        mem_matrixMask = clCreateSubBuffer(extra_mask->data_device, CL_MEM_READ_WRITE, CL_BUFFER_CREATE_TYPE_REGION, &region, &err);
+        CL_CHECK(err);
+
+        {
+            // create padded mask to contain all data
+            mem_matrixMask_padded = clCreateBuffer(context, CL_MEM_ALLOC_HOST_PTR, mask_nb_padded, NULL, &err);
+            CL_CHECK(err);
+
+            // pass extra_mask->data_device, mem_matrixMask to kernel for copying/padding
+            mask_nb1_padded = (cl_ulong)n_kv_padded * sizeof(cl_half);
+            mask_nb2_padded = mask_nb1_padded * (cl_ulong)mask->ne[1];
+            mask_nb3_padded = mask_nb2_padded * (cl_ulong)mask->ne[2];
+
+            cl_kernel repack_mask = backend_ctx->fa.kernel_repack_mask_for_wmm;
+            CL_CHECK(clSetKernelArg(repack_mask, 0, sizeof(cl_mem),   &mem_matrixMask));
+            CL_CHECK(clSetKernelArg(repack_mask, 1, sizeof(cl_ulong), &mask_nb1));
+            CL_CHECK(clSetKernelArg(repack_mask, 2, sizeof(cl_ulong), &mask_nb2));
+            CL_CHECK(clSetKernelArg(repack_mask, 3, sizeof(cl_ulong), &mask_nb3));
+            CL_CHECK(clSetKernelArg(repack_mask, 4, sizeof(int),      &mask_ne2));
+            CL_CHECK(clSetKernelArg(repack_mask, 5, sizeof(cl_mem),   &mem_matrixMask_padded));
+            CL_CHECK(clSetKernelArg(repack_mask, 6, sizeof(cl_ulong), &mask_nb1_padded));
+            CL_CHECK(clSetKernelArg(repack_mask, 7, sizeof(cl_ulong), &mask_nb2_padded));
+            CL_CHECK(clSetKernelArg(repack_mask, 8, sizeof(cl_ulong), &mask_nb3_padded));
+
+            size_t repack_mask_gws[3] = {(size_t)n_kv, (size_t)mask->ne[1], (size_t)mask_ne2 * (size_t)mask->ne[3]};
+            backend_ctx->enqueue_ndrange_kernel(repack_mask, 3, repack_mask_gws, NULL, dst);
+        }
+
+        // use image 1d buffer for matrix Mask (padded row stride)
+        cl_image_format img_fmt_mask_1d = { CL_RGBA, CL_HALF_FLOAT};
+        cl_image_desc img_desc_mask_1d;
+        memset(&img_desc_mask_1d, 0, sizeof(img_desc_mask_1d));
+        img_desc_mask_1d.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER;
+        img_desc_mask_1d.image_width = mask_nb_padded / 2 / 4;
+        img_desc_mask_1d.buffer = mem_matrixMask_padded;
+        mem_tex_matrixMask_1dbuf = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt_mask_1d, &img_desc_mask_1d, NULL, &err);
+        CL_CHECK(err);
+    }
+
+    // WMM QK uses repacked 3D images.
+    // Q image: rows, heads, packed depth.
+    cl_image_format img_fmt_3d = { CL_RGBA, CL_HALF_FLOAT };
+    cl_image_desc   img_desc_3d;
+
+    memset(&img_desc_3d, 0, sizeof(img_desc_3d));
+    img_desc_3d.image_type   = CL_MEM_OBJECT_IMAGE3D;
+    img_desc_3d.image_width  = (size_t)n_q;
+    img_desc_3d.image_height = (size_t)n_batch * (size_t)n_head;
+    img_desc_3d.image_depth  = (size_t)d_head_q / 4;
+    cl_mem img_q_wmm = NULL;
+    img_q_wmm = clCreateImage(context, CL_MEM_READ_WRITE, &img_fmt_3d, &img_desc_3d, NULL, &err);
+    CL_CHECK(err);
+
+    {
+        cl_kernel repack_q = backend_ctx->fa.kernel_repack_q_for_wmm;
+        CL_CHECK(clSetKernelArg(repack_q, 0, sizeof(cl_mem),   &mem_matrixQ));
+        CL_CHECK(clSetKernelArg(repack_q, 1, sizeof(cl_ulong), &q_nb1));
+        CL_CHECK(clSetKernelArg(repack_q, 2, sizeof(cl_ulong), &q_nb2));
+        CL_CHECK(clSetKernelArg(repack_q, 3, sizeof(cl_ulong), &q_nb3));
+        CL_CHECK(clSetKernelArg(repack_q, 4, sizeof(int),      &n_head));
+        CL_CHECK(clSetKernelArg(repack_q, 5, sizeof(cl_mem),   &img_q_wmm));
+
+        size_t repack_q_gws[3] = {(size_t)d_head_q / 4, (size_t)n_q, (size_t)n_batch * (size_t)n_head};
+        backend_ctx->enqueue_ndrange_kernel(repack_q, 3, repack_q_gws, NULL, dst);
+    }
+
+    // K image: columns, row groups, KV heads.
+    const size_t n_kv_row4 = ((size_t)n_kv + 3) / 4;
+
+    memset(&img_desc_3d, 0, sizeof(img_desc_3d));
+    img_desc_3d.image_type   = CL_MEM_OBJECT_IMAGE3D;
+    img_desc_3d.image_width  = (size_t)d_head_q;
+    img_desc_3d.image_height = n_kv_row4;
+    img_desc_3d.image_depth  = (size_t)n_batch * (size_t)n_head_kv;
+    cl_mem img_k_wmm = NULL;
+    img_k_wmm = clCreateImage(context, CL_MEM_READ_WRITE, &img_fmt_3d, &img_desc_3d, NULL, &err);
+    CL_CHECK(err);
+
+    {
+        cl_kernel repack_k = backend_ctx->fa.kernel_repack_k_for_wmm;
+        CL_CHECK(clSetKernelArg(repack_k, 0, sizeof(cl_mem),   &mem_matrixK));
+        CL_CHECK(clSetKernelArg(repack_k, 1, sizeof(cl_ulong), &k_nb1));
+        CL_CHECK(clSetKernelArg(repack_k, 2, sizeof(cl_ulong), &k_nb2));
+        CL_CHECK(clSetKernelArg(repack_k, 3, sizeof(cl_ulong), &k_nb3));
+        CL_CHECK(clSetKernelArg(repack_k, 4, sizeof(int),      &n_head_kv));
+        CL_CHECK(clSetKernelArg(repack_k, 5, sizeof(int),      &n_kv));
+        CL_CHECK(clSetKernelArg(repack_k, 6, sizeof(cl_mem),   &img_k_wmm));
+
+        size_t repack_k_gws[3] = {(size_t)d_head_q, n_kv_row4, (size_t)n_batch * (size_t)n_head_kv};
+        backend_ctx->enqueue_ndrange_kernel(repack_k, 3, repack_k_gws, NULL, dst);
+    }
+
+    // V image: kv-rows (contracted), packed head-dim groups, KV heads.
+    memset(&img_desc_3d, 0, sizeof(img_desc_3d));
+    img_desc_3d.image_type   = CL_MEM_OBJECT_IMAGE3D;
+    img_desc_3d.image_width  = (size_t)n_kv;
+    img_desc_3d.image_height = (size_t)d_head_v / 4;
+    img_desc_3d.image_depth  = (size_t)n_batch * (size_t)n_head_kv;
+    cl_mem img_v_wmm = NULL;
+    img_v_wmm = clCreateImage(context, CL_MEM_READ_WRITE, &img_fmt_3d, &img_desc_3d, NULL, &err);
+    CL_CHECK(err);
+
+    {
+        cl_kernel repack_v = backend_ctx->fa.kernel_repack_v_for_wmm;
+        CL_CHECK(clSetKernelArg(repack_v, 0, sizeof(cl_mem),   &mem_matrixV));
+        CL_CHECK(clSetKernelArg(repack_v, 1, sizeof(cl_ulong), &v_nb1));
+        CL_CHECK(clSetKernelArg(repack_v, 2, sizeof(cl_ulong), &v_nb2));
+        CL_CHECK(clSetKernelArg(repack_v, 3, sizeof(cl_ulong), &v_nb3));
+        CL_CHECK(clSetKernelArg(repack_v, 4, sizeof(int),      &n_head_kv));
+        CL_CHECK(clSetKernelArg(repack_v, 5, sizeof(cl_mem),   &img_v_wmm));
+
+        size_t repack_v_gws[3] = {(size_t)d_head_v / 4, (size_t)n_kv, (size_t)n_batch * (size_t)n_head_kv};
+        backend_ctx->enqueue_ndrange_kernel(repack_v, 3, repack_v_gws, NULL, dst);
+    }
+
+    cl_int enable_mask = (extra_mask) ? 1 : 0;
+    mask_buffer = extra_mask ? mem_tex_matrixMask_1dbuf : mem_tex_mask_fallback_1dbuf;
+
+    cl_mem mem_sinksBuf = NULL;
+    cl_mem mem_tex_sinks_1dbuf = NULL;
+    cl_int enable_sinks = (sinks_buffer != NULL) ? 1 : 0;
+    if (enable_sinks) {
+        region.origin = offset_sinks;
+        region.size = ggml_nbytes(sinks);
+        mem_sinksBuf = clCreateSubBuffer(extra_sinks->data_device, CL_MEM_READ_ONLY, CL_BUFFER_CREATE_TYPE_REGION, &region, &err);
+        CL_CHECK(err);
+
+        cl_image_format img_fmt_sinks_1d = { CL_R, CL_FLOAT };
+        cl_image_desc img_desc_sinks_1d;
+        memset(&img_desc_sinks_1d, 0, sizeof(img_desc_sinks_1d));
+        img_desc_sinks_1d.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER;
+        img_desc_sinks_1d.image_width = (size_t)n_head;
+        img_desc_sinks_1d.buffer = mem_sinksBuf;
+        mem_tex_sinks_1dbuf = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt_sinks_1d, &img_desc_sinks_1d, NULL, &err);
+        CL_CHECK(err);
+    } else {
+        // The image obj cannot be null so we back with buffer of size 1 and use matrixK to back because it always exists
+        cl_image_format img_fmt_sinks_fallback = { CL_R, CL_FLOAT };
+        cl_image_desc img_desc_sinks_fallback;
+        memset(&img_desc_sinks_fallback, 0, sizeof(img_desc_sinks_fallback));
+        img_desc_sinks_fallback.image_type = CL_MEM_OBJECT_IMAGE1D_BUFFER;
+        img_desc_sinks_fallback.image_width = 1;
+        img_desc_sinks_fallback.buffer = mem_matrixK;
+        mem_tex_sinks_1dbuf = clCreateImage(context, CL_MEM_READ_ONLY, &img_fmt_sinks_fallback, &img_desc_sinks_fallback, NULL, &err);
+        CL_CHECK(err);
+    }
+
+    cl_uint arg = 0;
+
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_mem),   &mem_tex_matrixO_1dbuf));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(float),    &scale));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int),      &n_q));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int),      &n_kv));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int),      &is_causal));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int),      &n_head));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &q_nb1));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &q_nb2));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &q_nb3));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &k_nb1));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &k_nb2));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &k_nb3));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &v_nb1));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &v_nb2));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &v_nb3));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &o_nb1));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &o_nb2));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &o_nb3));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(float),    &max_bias));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(float),    &m0));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(float),    &m1));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int),      &n_head_log2_val));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(float),    &logit_softcap));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int),      &n_head_kv));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_mem),   &mask_buffer));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int),      &enable_mask));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &mask_nb1_padded));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &mask_nb2_padded));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_ulong), &mask_nb3_padded));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int),      &mask_ne2));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int),      &mask_ne3));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_mem),   &mem_tex_sinks_1dbuf));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int),      &enable_sinks));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_mem),   &img_q_wmm));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_mem),   &img_k_wmm));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(cl_mem),   &img_v_wmm));
+    CL_CHECK(clSetKernelArg(kernel, arg++, sizeof(int),      &d_head_q));
+
+    size_t global_work_size[3], local_work_size[3];
+
+    const int n_waves_v = d_head_q / 64;
+
+    local_work_size[0] = 64;
+    local_work_size[1] = n_waves_v;
+    local_work_size[2] = 1;
+
+    global_work_size[0] = 64;
+    global_work_size[1] = ((n_q + 64 - 1) / 64) * n_waves_v;
+    global_work_size[2] = n_batch * n_head;
+
+    backend_ctx->enqueue_ndrange_kernel(kernel, 3, global_work_size, local_work_size, dst);
+
+    CL_CHECK(clReleaseMemObject(mem_tex_matrixO_1dbuf));
+    CL_CHECK(clReleaseMemObject(img_q_wmm));
+    CL_CHECK(clReleaseMemObject(img_k_wmm));
+    CL_CHECK(clReleaseMemObject(img_v_wmm));
+
+    if (mem_tex_matrixMask_1dbuf) {
+        CL_CHECK(clReleaseMemObject(mem_tex_matrixMask_1dbuf));
+    }
+    if (mem_matrixMask) {
+        CL_CHECK(clReleaseMemObject(mem_matrixMask));
+    }
+    if (mem_matrixMask_padded) {
+        CL_CHECK(clReleaseMemObject(mem_matrixMask_padded));
+    }
+    if (mem_tex_sinks_1dbuf) {
+        CL_CHECK(clReleaseMemObject(mem_tex_sinks_1dbuf));
+    }
+    if (mem_sinksBuf) {
+        CL_CHECK(clReleaseMemObject(mem_sinksBuf));
+    }
+    CL_CHECK(clReleaseMemObject(mem_matrixQ));
+    CL_CHECK(clReleaseMemObject(mem_matrixK));
+    CL_CHECK(clReleaseMemObject(mem_matrixV));
+    CL_CHECK(clReleaseMemObject(mem_matrixO));
+}
+#endif // GGML_OPENCL_USE_ADRENO_KERNELS
+
 static void ggml_cl_flash_attn(ggml_backend_t backend, const ggml_tensor * q, const ggml_tensor * k, ggml_tensor * dst) {
     const ggml_tensor * v = dst->src[2];
     const ggml_tensor * mask = dst->src[3];
@@ -17253,6 +17728,14 @@ static void ggml_cl_flash_attn(ggml_backend_t backend, const ggml_tensor * q, co
     const bool is_q8_0 = q->type == GGML_TYPE_F32 && k->type == GGML_TYPE_Q8_0 && v->type == GGML_TYPE_Q8_0;
     const bool is_q4_0 = q->type == GGML_TYPE_F32 && k->type == GGML_TYPE_Q4_0 && v->type == GGML_TYPE_Q4_0;

+#ifdef GGML_OPENCL_USE_ADRENO_KERNELS
+    if (use_fa_bin_kernels_prefill(backend_ctx, q, k, v)) {
+        // We support the prefill path of flash attn with a specialized d_head = 64/128/256
+        ggml_cl_flash_attn_prefill_bin(backend, q, k, dst);
+        return;
+    }
+#endif
+
     if (is_f16) {
         ggml_opencl_ensure_fa_variant(backend_ctx, d_head_q, d_head_v, FA_VARIANT_F16);
     } else if (is_mixed) {
diff --git a/ggml/src/ggml-opencl/kernels/flash_attn_repack.cl b/ggml/src/ggml-opencl/kernels/flash_attn_repack.cl
new file mode 100644
index 000000000..db78d5634
--- /dev/null
+++ b/ggml/src/ggml-opencl/kernels/flash_attn_repack.cl
@@ -0,0 +1,92 @@
+#pragma OPENCL EXTENSION cl_khr_fp16 : enable
+
+__kernel void kernel_repack_mask_for_wmm(
+    const global half* mask_buf,
+    const ulong mask_nb1,
+    const ulong mask_nb2,
+    const ulong mask_nb3,
+    const int mask_ne2,
+    global half* mask_buf_padded,
+    const ulong mask_nb1_padded,
+    const ulong mask_nb2_padded,
+    const ulong mask_nb3_padded
+) {
+    int col   = get_global_id(0);   // 0 .. n_kv
+    int row   = get_global_id(1);   // 0 .. n_q
+    int slice = get_global_id(2);   // 0 .. (n_head * n_batch)
+
+    int head_idx  = slice % mask_ne2;
+    int batch_idx = slice / mask_ne2;
+
+    ulong src_off = (ulong)batch_idx * mask_nb3 + (ulong)head_idx * mask_nb2 + (ulong)row * mask_nb1;
+    ulong dst_off = (ulong)batch_idx * mask_nb3_padded + (ulong)head_idx * mask_nb2_padded + (ulong)row * mask_nb1_padded;
+
+    mask_buf_padded[dst_off / 2 + col] = mask_buf[src_off / 2 + col];
+}
+
+__kernel void kernel_repack_q_for_wmm(
+    const global float* q_buf,
+    const ulong q_nb1,
+    const ulong q_nb2,
+    const ulong q_nb3,
+    const int n_head,
+    __write_only image3d_t img_q_wmm
+) {
+    int k4    = get_global_id(0);
+    int row   = get_global_id(1);
+    int slice = get_global_id(2);
+    int batch_idx = slice / n_head;
+    int head_idx  = slice % n_head;
+
+
+    ulong elem_off = (batch_idx * q_nb3 + head_idx * q_nb2 + row * q_nb1) / 4 + (ulong)k4 * 4;
+    float4 v = vload4(elem_off / 4, q_buf);
+
+    write_imageh(img_q_wmm, (int4)(row, slice, k4, 0), convert_half4(v));
+}
+
+__kernel void kernel_repack_k_for_wmm(
+    const global half* k_buf,
+    const ulong k_nb1,
+    const ulong k_nb2,
+    const ulong k_nb3,
+    const int n_head_kv,
+    const int n_kv,
+    __write_only image3d_t img_k_wmm
+) {
+    int kk    = get_global_id(0);
+    int row4  = get_global_id(1);
+    int slice = get_global_id(2);
+    int batch_idx   = slice / n_head_kv;
+    int head_kv_idx = slice % n_head_kv;
+
+    ulong base = batch_idx * k_nb3 + head_kv_idx * k_nb2;
+    int row0 = row4 * 4;
+    half4 v;
+    v.x = (row0 + 0 < n_kv) ? k_buf[(base + (ulong)(row0 + 0) * k_nb1) / 2 + kk] : (half)0;
+    v.y = (row0 + 1 < n_kv) ? k_buf[(base + (ulong)(row0 + 1) * k_nb1) / 2 + kk] : (half)0;
+    v.z = (row0 + 2 < n_kv) ? k_buf[(base + (ulong)(row0 + 2) * k_nb1) / 2 + kk] : (half)0;
+    v.w = (row0 + 3 < n_kv) ? k_buf[(base + (ulong)(row0 + 3) * k_nb1) / 2 + kk] : (half)0;
+
+    write_imageh(img_k_wmm, (int4)(kk, row4, slice, 0), v);
+}
+
+__kernel void kernel_repack_v_for_wmm(
+      const global half* v_buf,
+      const ulong v_nb1,
+      const ulong v_nb2,
+      const ulong v_nb3,
+      const int n_head_kv,
+      __write_only image3d_t img_v_wmm
+) {
+    int hdim4   = get_global_id(0);   // now fastest — walks contiguous memory
+    int row     = get_global_id(1);
+    int slice   = get_global_id(2);
+    int batch_idx   = slice / n_head_kv;
+    int head_kv_idx = slice % n_head_kv;
+
+    ulong row_off = batch_idx * v_nb3 + head_kv_idx * v_nb2 + (ulong)row * v_nb1;
+    half4 v = vload4((row_off / 2 + (ulong)hdim4 * 4) / 4, v_buf);
+
+    write_imageh(img_v_wmm, (int4)(row, hdim4, slice, 0), v);
+}