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, ®ion, &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, ®ion, &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, ®ion, &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, ®ion, &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, ®ion, &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, ®ion, &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);
+}