Commit f2cc7282c for llama.cpp

commit f2cc7282ce0042641df55e0e563b75e824e2b647
Author: lhez <lih@qti.qualcomm.com>
Date:   Sun Oct 11 11:49:25 2026 -0700

    opencl: improve fa, allow dk512 for gemma-4, improve dk64 (#30266)

    * opencl: enable Gemma-4 E4B GPU decode

    Co-authored-by: Hongqiang Wang <wangh@qti.qualcomm.com>

    * opencl: extend Gemma-4 GPU decode to E2B

    Co-authored-by: Hongqiang Wang <wangh@qti.qualcomm.com>

    * opencl: optimize DK64 GQA8 decode

    Co-authored-by: Hongqiang Wang <wangh@qti.qualcomm.com>

    * opencl: optimize DK128 GQA4 decode

    Co-authored-by: Hongqiang Wang <wangh@qti.qualcomm.com>

    ---------

    Co-authored-by: Hongqiang Wang <wangh@qti.qualcomm.com>

diff --git a/ggml/src/ggml-opencl/ggml-opencl.cpp b/ggml/src/ggml-opencl/ggml-opencl.cpp
index 625b12f54..2abc0bc58 100644
--- a/ggml/src/ggml-opencl/ggml-opencl.cpp
+++ b/ggml/src/ggml-opencl/ggml-opencl.cpp
@@ -48,6 +48,7 @@ typedef const void * (*get_adreno_bin_kernel_func_t)(
 #include <mutex>
 #include <regex>
 #include <set>
+#include <tuple>
 #include <unordered_set>

 #undef MIN
@@ -494,6 +495,13 @@ struct ggml_opencl_fa_kernels {
     std::map<std::pair<int, int>, cl_kernel> f32_f16_q1_split;       // flash-decoding K-split
     // vec decode
     std::map<std::pair<int, int>, cl_kernel> f32_f16_q1_vec;
+    bool f32_f16_vec_512_attempted = false;
+    std::map<std::tuple<int, int, int>, cl_kernel> f32_f16_mq_decode;
+    std::map<std::tuple<int, int, int>, size_t>    f32_f16_mq_decode_wg;
+    std::map<std::tuple<int, int, int>, int>       f32_f16_mq_decode_hs;
+    std::set<std::tuple<int, int, int>>           f32_f16_mq_decode_attempted;
+    ggml_cl_buffer fd_partial;
+    cl_uint compute_units = 0;
     // kv-head-coalesced vec decode
     std::map<std::pair<int, int>, cl_kernel> f32_f16_q1_vec_mq;
     // kv-head-coalesced + flash-decoding split
@@ -1289,6 +1297,11 @@ struct ggml_backend_opencl_context {

         ref_count--;
         if (ref_count == 0) {
+            if (fa.fd_partial.buffer) {
+                CL_CHECK(clReleaseMemObject(fa.fd_partial.buffer));
+                fa.fd_partial.buffer = nullptr;
+                fa.fd_partial.size = 0;
+            }
 #ifdef GGML_OPENCL_PROFILING
             flush_profiling_batch();
             write_profiling_info();
@@ -1473,14 +1486,7 @@ static bool use_adreno_bin_kernels(ggml_backend_opencl_context * backend_ctx) {
 #endif // GGML_OPENCL_USE_ADRENO_BIN_KERNELS
 }

-static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) {
-    if (backend_ctx->kernels_loaded) {
-        return;
-    }
-
-    cl_int err;
-
-    // compiler options for general kernels
+static std::string ggml_opencl_make_compile_opts(const ggml_backend_opencl_context * backend_ctx) {
     auto opencl_c_std =
         std::string("CL") + std::to_string(backend_ctx->opencl_c_version.major) + "." + std::to_string(backend_ctx->opencl_c_version.minor);
     std::string compile_opts = std::string("-cl-std=") + opencl_c_std +
@@ -1491,7 +1497,18 @@ static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) {
         compile_opts += " -qcom-enable-large-buffer ";
     }

-    backend_ctx->kernel_compile_opts = compile_opts;
+    return compile_opts;
+}
+
+static void load_cl_kernels(ggml_backend_opencl_context *backend_ctx) {
+    if (backend_ctx->kernels_loaded) {
+        return;
+    }
+
+    cl_int err;
+    const std::string & compile_opts = backend_ctx->kernel_compile_opts;
+    const std::string opencl_c_std = "CL" + std::to_string(backend_ctx->opencl_c_version.major) +
+                                   "." + std::to_string(backend_ctx->opencl_c_version.minor);

     GGML_LOG_INFO("ggml_opencl: loading OpenCL kernels");

@@ -5606,6 +5623,127 @@ static void ggml_opencl_ensure_fa_pre_kernels(ggml_backend_opencl_context * back
     clReleaseProgram(prog_pre_f16);
 }

+static bool ggml_opencl_ensure_fa_f32_f16_vec_512(ggml_backend_opencl_context * backend_ctx) {
+    const std::pair<int, int> key = {512, 512};
+    auto & fa = backend_ctx->fa;
+    if (fa.f32_f16_q1_vec.count(key) > 0) {
+        return true;
+    }
+    if (fa.f32_f16_vec_512_attempted || backend_ctx->kernel_compile_opts.empty()) {
+        return false;
+    }
+    fa.f32_f16_vec_512_attempted = true;
+
+    const ggml_opencl_fa_dim * cfg = nullptr;
+    for (const auto & d : g_opencl_fa_dims) {
+        if (d.dk == 512 && d.dv == 512) {
+            cfg = &d;
+            break;
+        }
+    }
+    if (cfg == nullptr) {
+        return false;
+    }
+
+    // Compile only vec decode and merge to stay within the Adreno compiler's memory limit.
+    const std::string opts = ggml_opencl_fa_compile_opts(backend_ctx, cfg, FA_VARIANT_F32_F16) +
+                             " -D FA_DECODE_ONLY -D FA_VEC_ONLY";
+    cl_program prog = build_program_from_source_ex(
+        backend_ctx->context, backend_ctx->device,
+        ggml_opencl_fa_kernel_src(FA_VARIANT_F32_F16).c_str(), opts,
+        /*fatal=*/false, "fa f32_f16 decode512 vec", backend_ctx->queue);
+    if (!prog) {
+        return false;
+    }
+    cl_int err;
+    cl_kernel vec = clCreateKernel(prog, "flash_attn_f32_f16_q1_vec", &err);
+    if (err != CL_SUCCESS) {
+        clReleaseProgram(prog);
+        return false;
+    }
+    cl_kernel merge = clCreateKernel(prog, "flash_attn_f32_merge", &err);
+    clReleaseProgram(prog);
+    if (err != CL_SUCCESS) {
+        clReleaseKernel(vec);
+        return false;
+    }
+    if (!ggml_opencl_fa_kernel_fits_wg(backend_ctx, vec, 256, "flash_attn_f32_f16_q1_vec", 512, 512) ||
+        !ggml_opencl_fa_kernel_fits_wg(backend_ctx, merge, 128, "flash_attn_f32_merge", 512, 512)) {
+        clReleaseKernel(vec);
+        clReleaseKernel(merge);
+        return false;
+    }
+    fa.f32_f16_q1_vec[key] = vec;
+    if (fa.f32_merge.count(key) > 0) {
+        clReleaseKernel(merge);
+    } else {
+        fa.f32_merge[key] = merge;
+    }
+    return true;
+}
+
+static void ggml_opencl_ensure_fa_f32_f16_mq_decode(ggml_backend_opencl_context * backend_ctx, int dk, int dv, int gqa) {
+    const std::tuple<int, int, int> key = {dk, dv, gqa};
+    auto & fa = backend_ctx->fa;
+    if (fa.f32_f16_mq_decode.count(key) > 0 || fa.f32_f16_mq_decode_attempted.count(key) > 0 ||
+        backend_ctx->kernel_compile_opts.empty()) {
+        return;
+    }
+    if (gqa == 4 && dk != 128 && fa.f32_f16_q1_vec_mq_split.count({dk, dv}) > 0) {
+        fa.f32_f16_mq_decode[key] = fa.f32_f16_q1_vec_mq_split.at({dk, dv});
+        fa.f32_f16_mq_decode_wg[key] = 256;
+        fa.f32_f16_mq_decode_hs[key] = 1;
+        return;
+    }
+    fa.f32_f16_mq_decode_attempted.insert(key);
+    const ggml_opencl_fa_dim * cfg = nullptr;
+    for (const auto & d : g_opencl_fa_dims) {
+        if (d.dk == dk && d.dv == dv) {
+            cfg = &d;
+            break;
+        }
+    }
+    if (cfg == nullptr) {
+        return;
+    }
+
+    const bool cluster = dk == 64 || dk == 128;
+    const int head_sub = cluster ? 2 : (gqa == 8 ? (dk == 512 ? 4 : 2) : 1);
+    const int nsg_max = dk == 64 ? 1 : (dk == 128 || (gqa == 8 && dk == 256) ? 2 : 4);
+    const char * kernel_name = cluster ? "flash_attn_f32_f16_q1_vec_mq_split_c8" : "flash_attn_f32_f16_q1_vec_mq_split";
+    const std::string src = ggml_opencl_fa_kernel_src(FA_VARIANT_F32_F16);
+    const std::string opts = ggml_opencl_fa_compile_opts(backend_ctx, cfg, FA_VARIANT_F32_F16) +
+                             " -D FA_MQ_ONLY -D MQ_GQA=" + std::to_string(gqa / head_sub) +
+                             " -D FA_HEAD_SUB=" + std::to_string(head_sub) +
+                             (cluster ? " -D MQ_NSG=" + std::to_string(nsg_max) + " -D FA_CL_C=16" : " -D FA_MQ_SPLIT_ONLY") +
+                             (dk == 64 ? " -D FA_CL_MHRED -D FA_CL_MASK_BCAST" : "") +
+                             (gqa == 8 && dk == 256 ? " -D FA_Q1_Q_REG" : "");
+    for (int nsg = nsg_max; nsg >= 1; nsg /= 2) {
+        const size_t wg = 64 * nsg;
+        cl_program prog = build_program_from_source_ex(
+            backend_ctx->context, backend_ctx->device, src.c_str(),
+            opts + " -D MQ_NSG_SPLIT=" + std::to_string(nsg),
+            /*fatal=*/false, "fa f32_f16 mq decode", backend_ctx->queue);
+        if (!prog) {
+            continue;
+        }
+        cl_int err;
+        cl_kernel kernel = clCreateKernel(prog, kernel_name, &err);
+        clReleaseProgram(prog);
+        if (err != CL_SUCCESS) {
+            continue;
+        }
+        if (!ggml_opencl_fa_kernel_fits_wg(backend_ctx, kernel, wg, kernel_name, dk, dv)) {
+            clReleaseKernel(kernel);
+            continue;
+        }
+        fa.f32_f16_mq_decode[key] = kernel;
+        fa.f32_f16_mq_decode_wg[key] = wg;
+        fa.f32_f16_mq_decode_hs[key] = head_sub;
+        return;
+    }
+}
+
 // DK=512 prefill BM-tile
 static bool ggml_opencl_ensure_fa_f32_f16_prefill_512(ggml_backend_opencl_context * backend_ctx, bool split) {
     const int dk = 512, dv = 512;
@@ -6834,6 +6972,10 @@ static ggml_backend_opencl_context * ggml_cl_init(ggml_backend_dev_t dev) {
     backend_ctx->adreno_use_large_buffer = getenv("GGML_OPENCL_ADRENO_USE_LARGE_BUFFER") != nullptr &&
                                            backend_ctx->gpu_family == GPU_FAMILY::ADRENO;

+    backend_ctx->kernel_compile_opts = ggml_opencl_make_compile_opts(backend_ctx.get());
+    CL_CHECK(clGetDeviceInfo(device, CL_DEVICE_MAX_COMPUTE_UNITS,
+                            sizeof(backend_ctx->fa.compute_units), &backend_ctx->fa.compute_units, NULL));
+
     // ragged moe, unspecified or non-zero means enabled, set to 0 to disable
     static const char * ragged_fp16_env = getenv("GGML_OPENCL_MOE_RAGGED_FP16");
     backend_ctx->adreno_use_moe_ragged = (ragged_fp16_env == NULL) ? 1 : (atoi(ragged_fp16_env) != 0);
@@ -9347,10 +9489,13 @@ static bool ggml_opencl_supports_op(ggml_backend_dev_t dev, const struct ggml_te
                     return false;
                 }
                 if (q->ne[1] == 1) {
-                    // DK=512 decode is bandwidth-bound and slower on the GPU
-                    // than on the CPU; decline it here so it runs on the CPU.
-                    // Prefill (n_q > 1) stays on the GPU.
-                    return false;
+                    const char * decode_env = getenv("GGML_OPENCL_FA_DK512_DECODE");
+                    if ((decode_env && decode_env[0] == '0') ||
+                        backend_ctx->gpu_family != ADRENO || k->ne[2] <= 0 ||
+                        (q->ne[2] / k->ne[2] != 4 && q->ne[2] / k->ne[2] != 8) || q->ne[2] % k->ne[2] != 0 ||
+                        !ggml_opencl_ensure_fa_f32_f16_vec_512(backend_ctx)) {
+                        return false;
+                    }
                 } else {
                     // prefill, BM-tile in its own FA_PREFILL_ONLY program
                     if (!ggml_opencl_ensure_fa_f32_f16_prefill_512(backend_ctx, /*split=*/false)) {
@@ -17938,10 +18083,7 @@ static void ggml_cl_flash_attn(ggml_backend_t backend, const ggml_tensor * q, co
     }
 #endif

-    // DK=512 (Gemma-4 global layers) runs decode-only (q1 / q1_split) on
-    // Adreno - it never uses the BM-tile path, and the prepass + split-tile
-    // programs OOM the compiler at DK=512; supports_op only admits
-    // n_q==1 here and prefill goes to CPU
+    // Compile DK512 decode separately from the prefill programs.
     const bool fa_decode_only_512 = (d_head_q == 512);

     // per-variant lazy compile for this (dk, dv)
@@ -17970,7 +18112,11 @@ static void ggml_cl_flash_attn(ggml_backend_t backend, const ggml_tensor * q, co
     if (is_f16) {
         ggml_opencl_ensure_fa_variant(backend_ctx, d_head_q, d_head_v, FA_VARIANT_F16);
     } else if (is_mixed) {
-        ggml_opencl_ensure_fa_variant(backend_ctx, d_head_q, d_head_v, FA_VARIANT_F32_F16);
+        if (fa_decode_only_512 && n_q == 1) {
+            GGML_ASSERT(ggml_opencl_ensure_fa_f32_f16_vec_512(backend_ctx));
+        } else {
+            ggml_opencl_ensure_fa_variant(backend_ctx, d_head_q, d_head_v, FA_VARIANT_F32_F16);
+        }
         if (fa_decode_only_512) {
             // DK=512: the BM-tile prefill kernels are specifically compiled from
             // FA_PREFILL_ONLY
@@ -18004,6 +18150,16 @@ static void ggml_cl_flash_attn(ggml_backend_t backend, const ggml_tensor * q, co
     }

     const std::pair<int, int> dk_dv = {d_head_q, d_head_v};
+    const int mq_gqa = n_head_kv > 0 ? n_head / n_head_kv : 0;
+    const std::tuple<int, int, int> mq_decode_key = {d_head_q, d_head_v, mq_gqa};
+    const bool mq_decode_shape = backend_ctx->gpu_family == ADRENO && is_mixed && n_q == 1 &&
+                                 d_head_q == d_head_v && n_head_kv > 0 && n_head % n_head_kv == 0 &&
+                                 ((((d_head_q == 64 && mq_gqa == 8) || (d_head_q == 128 && mq_gqa == 4)) &&
+                                   backend_ctx->has_subgroup_shuffle) ||
+                                  ((d_head_q == 256 || d_head_q == 512) && (mq_gqa == 4 || mq_gqa == 8)));
+    if (mq_decode_shape && n_kv >= 32) {
+        ggml_opencl_ensure_fa_f32_f16_mq_decode(backend_ctx, d_head_q, d_head_v, mq_gqa);
+    }
     const bool use_native_q8_0_q1 = is_q8_0 && n_q == 1 &&
                                     backend_ctx->fa.f32_q8_0_q1.count(dk_dv) > 0;
     // Native q8_0 prefill — reads q8_0 directly, wg_size = cfg->bm.
@@ -18262,6 +18418,8 @@ static void ggml_cl_flash_attn(ggml_backend_t backend, const ggml_tensor * q, co
     const int fd_max_n_q = (d_head_q <= FD_MAX_DK_MULTI) ? FD_MAX_N_Q_MULTI : 1;
     cl_kernel fd_k_split = NULL;
     bool use_fd_mq = false;
+    bool use_fd_mq_decode = false;
+    int fd_head_sub = 1;
     size_t fd_mq_wg = 256;  // MQ_GQA=4 kernel: Q1_WG_SIZE(64) * MQ_NSG_SPLIT(4)
     bool use_fa_k_img = false;  // K bound as image1d_buffer_t instead of (buf, offset)

@@ -18294,9 +18452,15 @@ static void ggml_cl_flash_attn(ggml_backend_t backend, const ggml_tensor * q, co
         if (mq_enabled && mq_kv_ok && nq_in_vec_range && !is_causal &&
             backend_ctx->gpu_family != INTEL &&
             !use_local_tile &&
-            n_kv >= FD_MIN_N_KV &&
+            n_kv >= (mq_decode_shape ? 32 : FD_MIN_N_KV) &&
             backend_ctx->fa.f32_merge.count(dk_dv) > 0) {
-            if (nq1_only && lmq_on && is_mixed && d_head_q == 128 && d_head_v == 128 &&
+            if (mq_decode_shape && backend_ctx->fa.f32_f16_mq_decode.count(mq_decode_key) > 0) {
+                fd_k_split = backend_ctx->fa.f32_f16_mq_decode.at(mq_decode_key);
+                fd_mq_wg = backend_ctx->fa.f32_f16_mq_decode_wg.at(mq_decode_key);
+                fd_head_sub = backend_ctx->fa.f32_f16_mq_decode_hs.at(mq_decode_key);
+                use_fd_mq = true;
+                use_fd_mq_decode = true;
+            } else if (nq1_only && lmq_on && is_mixed && d_head_q == 128 && d_head_v == 128 &&
                 gqa_ratio_dispatch == 8 &&
                 backend_ctx->fa.f32_f16_q1_local_mq_split_g8.count(dk_dv) > 0) {
                 fd_k_split = backend_ctx->fa.f32_f16_q1_local_mq_split_g8.at(dk_dv);
@@ -18575,6 +18739,14 @@ static void ggml_cl_flash_attn(ggml_backend_t backend, const ggml_tensor * q, co
         int n_splits = (n_kv + fd_kv_per_split - 1) / fd_kv_per_split;
         if (n_splits < FD_MIN_SPLITS) { n_splits = FD_MIN_SPLITS; }
         if (n_splits > fd_max_splits) { n_splits = fd_max_splits; }
+        if (use_fd_mq_decode) {
+            const size_t wg_per_split = (size_t) n_head_kv * n_batch;
+            const size_t wg_target = 4 * (size_t) backend_ctx->fa.compute_units;
+            while (wg_per_split * n_splits < wg_target && n_splits < fd_max_splits &&
+                   n_kv / (n_splits + 1) >= 32) {
+                n_splits++;
+            }
+        }
         const int kv_per_split = (n_kv + n_splits - 1) / n_splits;

         const int fa_partial_floats = 2 + d_head_v;
@@ -18582,15 +18754,26 @@ static void ggml_cl_flash_attn(ggml_backend_t backend, const ggml_tensor * q, co
             (size_t) n_batch * n_head * n_q * n_splits * fa_partial_floats * sizeof(float);

         ggml_cl_flash_attn_temp_buffer temp_partial;
+        cl_mem partial_buffer;
         cl_int err;
-        temp_partial.data = clCreateBuffer(backend_ctx->context, CL_MEM_READ_WRITE,
-                                           partial_size_bytes, NULL, &err);
-        if (err != CL_SUCCESS) {
-            CL_CHECK(clFinish(backend_ctx->queue));
+        if (use_fd_mq_decode) {
+            auto & pool = backend_ctx->fa.fd_partial;
+            if (partial_size_bytes > pool.size) {
+                CL_CHECK(clFinish(backend_ctx->queue));
+                pool.allocate(backend_ctx->context, partial_size_bytes);
+            }
+            partial_buffer = pool.buffer;
+        } else {
             temp_partial.data = clCreateBuffer(backend_ctx->context, CL_MEM_READ_WRITE,
                                                partial_size_bytes, NULL, &err);
+            if (err != CL_SUCCESS) {
+                CL_CHECK(clFinish(backend_ctx->queue));
+                temp_partial.data = clCreateBuffer(backend_ctx->context, CL_MEM_READ_WRITE,
+                                                   partial_size_bytes, NULL, &err);
+            }
+            CL_CHECK(err);
+            partial_buffer = temp_partial.data;
         }
-        CL_CHECK(err);

         cl_kernel k_split = fd_k_split;
         int argi = 0;
@@ -18658,7 +18841,7 @@ static void ggml_cl_flash_attn(ggml_backend_t backend, const ggml_tensor * q, co
         CL_CHECK(clSetKernelArg(k_split, argi++, sizeof(cl_ulong), &mask_nb3));
         CL_CHECK(clSetKernelArg(k_split, argi++, sizeof(int),      &mask_ne2));
         CL_CHECK(clSetKernelArg(k_split, argi++, sizeof(int),      &mask_ne3));
-        CL_CHECK(clSetKernelArg(k_split, argi++, sizeof(cl_mem),   &temp_partial.data));
+        CL_CHECK(clSetKernelArg(k_split, argi++, sizeof(cl_mem),   &partial_buffer));
         CL_CHECK(clSetKernelArg(k_split, argi++, sizeof(int),      &n_splits));
         CL_CHECK(clSetKernelArg(k_split, argi++, sizeof(int),      &kv_per_split));

@@ -18666,7 +18849,7 @@ static void ggml_cl_flash_attn(ggml_backend_t backend, const ggml_tensor * q, co
         // matches Q1_WG_SIZE * NSG (MQ_GQA=4 -> 256; MQ_GQA=8 -> 192)
         const size_t fd_wg = use_fd_mq ? fd_mq_wg : 64;
         const size_t fd_head_dim = use_fd_mq
-            ? (size_t)(n_head_kv * n_batch)
+            ? (size_t)(n_head_kv * fd_head_sub * n_batch)
             : (size_t)(n_head     * n_batch);
         size_t fd_lws[3] = { fd_wg, 1, 1 };
         // gid(2) packs q_idx * n_splits + split_idx.
@@ -18675,7 +18858,7 @@ static void ggml_cl_flash_attn(ggml_backend_t backend, const ggml_tensor * q, co

         cl_kernel k_merge = backend_ctx->fa.f32_merge.at(dk_dv);
         argi = 0;
-        CL_CHECK(clSetKernelArg(k_merge, argi++, sizeof(cl_mem),   &temp_partial.data));
+        CL_CHECK(clSetKernelArg(k_merge, argi++, sizeof(cl_mem),   &partial_buffer));
         CL_CHECK(clSetKernelArg(k_merge, argi++, sizeof(cl_mem),   &extra_o->data_device));
         CL_CHECK(clSetKernelArg(k_merge, argi++, sizeof(cl_ulong), &offset_o));
         CL_CHECK(clSetKernelArg(k_merge, argi++, sizeof(int),      &n_head));
diff --git a/ggml/src/ggml-opencl/kernels/flash_attn_f32_f16.cl b/ggml/src/ggml-opencl/kernels/flash_attn_f32_f16.cl
index bf7695a2c..27553c729 100644
--- a/ggml/src/ggml-opencl/kernels/flash_attn_f32_f16.cl
+++ b/ggml/src/ggml-opencl/kernels/flash_attn_f32_f16.cl
@@ -665,7 +665,7 @@ __kernel void FA_TILE_NAME(

 // allow bypassing decode kernels to avoid compiler crash for DK=512 on Adreno GPUs
 #ifndef FA_PREFILL_ONLY
-#ifndef FA_MQ_ONLY  // q1 excluded from the MQ-only (g8) program
+#if !defined(FA_MQ_ONLY) && !defined(FA_VEC_ONLY)
 REQD_FA_SG
 __kernel void flash_attn_f32_f16_q1(
     const global void * q_void, ulong q_offset,
@@ -932,14 +932,14 @@ __kernel void flash_attn_f32_f16_q1_vec(
         }
         ACC_TYPE dot_partial = dot4.s0 + dot4.s1 + dot4.s2 + dot4.s3;
         ACC_TYPE score = sub_group_reduce_add(dot_partial) * scale;
+        if (logit_softcap > 0.0f) {
+            score = logit_softcap * tanh(score / logit_softcap);
+        }

         if (mask_base != NULL) {
             const global MASK_DATA_TYPE * mask_ptr = (const global MASK_DATA_TYPE *) mask_base;
             score += slope * (ACC_TYPE) mask_ptr[k_idx];
         }
-        if (logit_softcap > 0.0f) {
-            score = logit_softcap * tanh(score / logit_softcap);
-        }

         // FA-2 online update. All threads in the subgroup see the same score,
         // so m_i and l_i evolve identically across lanes within the subgroup.
@@ -1385,6 +1385,7 @@ __kernel void flash_attn_f32_f16_q1_local_mq_split(
 #endif
 #define MQ_WG_SIZE (Q1_WG_SIZE * MQ_NSG)

+#ifndef FA_MQ_SPLIT_ONLY
 REQD_SUBGROUP_SIZE_64
 __kernel void flash_attn_f32_f16_q1_vec_mq(
     const global void * q_void, ulong q_offset,
@@ -1606,6 +1607,8 @@ __kernel void flash_attn_f32_f16_q1_vec_mq(
     }
 }

+#endif  // !FA_MQ_SPLIT_ONLY
+
 #ifndef MQ_NSG_SPLIT
 #define MQ_NSG_SPLIT 4
 #endif
@@ -1615,6 +1618,10 @@ __kernel void flash_attn_f32_f16_q1_vec_mq(
 #define FA_PARTIAL_FLOATS (2 + DV)
 #endif

+#ifndef FA_HEAD_SUB
+#define FA_HEAD_SUB 1
+#endif
+
 REQD_SUBGROUP_SIZE_64
 __kernel void flash_attn_f32_f16_q1_vec_mq_split(
     const global void * q_void, ulong q_offset,
@@ -1652,8 +1659,12 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split(
     const int split_idx        = split_q_idx % n_splits;
     const int q_idx            = split_q_idx / n_splits;

-    const int batch_idx   = kvhead_batch_idx / n_head_kv;
-    const int head_kv_idx = kvhead_batch_idx % n_head_kv;
+    const int hgroups     = n_head_kv * FA_HEAD_SUB;
+    const int batch_idx   = kvhead_batch_idx / hgroups;
+    const int hg          = kvhead_batch_idx % hgroups;
+    const int head_kv_idx = hg / FA_HEAD_SUB;
+    const int head_sub    = hg % FA_HEAD_SUB;
+#define FA_MQS_HEAD_IDX(h) (head_kv_idx * (MQ_GQA * FA_HEAD_SUB) + head_sub * MQ_GQA + (h))

     const int kv_start = split_idx * kv_per_split;
     const int kv_end   = min(kv_start + kv_per_split, n_kv);
@@ -1666,7 +1677,7 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split(
         if (tid == 0) {
             #pragma unroll
             for (int h = 0; h < MQ_GQA; ++h) {
-                const int head_idx = head_kv_idx * MQ_GQA + h;
+                const int head_idx = FA_MQS_HEAD_IDX(h);
                 const ulong rec_idx = ((((ulong) batch_idx * n_head + head_idx) * n_q + q_idx)
                                        * n_splits + split_idx);
                 global float * rec = partial_void + rec_idx * record_stride;
@@ -1681,22 +1692,33 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split(
     const global char * k_base = (const global char *) k_void + k_offset;
     const global char * v_base = (const global char *) v_void + v_offset;

+#ifdef FA_Q1_Q_REG
+    ACC_TYPE4 q_reg[MQ_GQA];
+    #pragma unroll
+    for (int h = 0; h < MQ_GQA; ++h) {
+        const int head_idx = FA_MQS_HEAD_IDX(h);
+        const ulong q_row_offset = batch_idx * q_nb3 + head_idx * q_nb2 + (ulong) q_idx * q_nb1;
+        const global Q_DATA_TYPE4 * q_ptr = (const global Q_DATA_TYPE4 *) (q_base + q_row_offset);
+        q_reg[h] = (tid_sg < DK_VEC) ? CONVERT_Q_ACC4(q_ptr[tid_sg]) : (ACC_TYPE4)(0.0f);
+    }
+#else
     // stage MQ_GQA Q rows in __local once (uniform across WG)
     __local ACC_TYPE4 q_shared[MQ_GQA * DK_VEC];
     for (int i = tid; i < MQ_GQA * DK_VEC; i += MQ_SPLIT_WG_SIZE) {
         const int h        = i / DK_VEC;
         const int k        = i % DK_VEC;
-        const int head_idx = head_kv_idx * MQ_GQA + h;
+        const int head_idx = FA_MQS_HEAD_IDX(h);
         const ulong q_row_offset = batch_idx * q_nb3 + head_idx * q_nb2 + (ulong) q_idx * q_nb1;
         const global Q_DATA_TYPE4 * q_ptr = (const global Q_DATA_TYPE4 *) (q_base + q_row_offset);
         q_shared[h * DK_VEC + k] = CONVERT_Q_ACC4(q_ptr[k]);
     }
     barrier(CLK_LOCAL_MEM_FENCE);
+#endif

     float slope[MQ_GQA];
     #pragma unroll
     for (int h = 0; h < MQ_GQA; ++h) {
-        slope[h] = get_alibi_slope(max_bias, head_kv_idx * MQ_GQA + h, n_head_log2, m0, m1);
+        slope[h] = get_alibi_slope(max_bias, FA_MQS_HEAD_IDX(h), n_head_log2, m0, m1);
     }

     const global char * mask_base[MQ_GQA];
@@ -1707,7 +1729,7 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split(
                                           (ulong) q_idx * mask_nb1;
         #pragma unroll
         for (int h = 0; h < MQ_GQA; ++h) {
-            const int head_idx      = head_kv_idx * MQ_GQA + h;
+            const int head_idx      = FA_MQS_HEAD_IDX(h);
             const int mask_head_idx = head_idx % mask_ne2;
             mask_base[h] = mask_base_b + mask_head_idx * mask_nb2;
         }
@@ -1742,6 +1764,15 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split(
         ACC_TYPE4 dot4[MQ_GQA];
         #pragma unroll
         for (int h = 0; h < MQ_GQA; ++h) dot4[h] = (ACC_TYPE4)(0.0f);
+#ifdef FA_Q1_Q_REG
+        if (tid_sg < DK_VEC) {
+            const ACC_TYPE4 k_vec = CONVERT_KV_ACC4(k_ptr[tid_sg]);
+            #pragma unroll
+            for (int h = 0; h < MQ_GQA; ++h) {
+                dot4[h] = mad(q_reg[h], k_vec, dot4[h]);
+            }
+        }
+#else
         for (int k = tid_sg; k < DK_VEC; k += Q1_WG_SIZE) {
             const ACC_TYPE4 k_vec = CONVERT_KV_ACC4(k_ptr[k]);
             #pragma unroll
@@ -1749,19 +1780,20 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split(
                 dot4[h] = mad(q_shared[h * DK_VEC + k], k_vec, dot4[h]);
             }
         }
+#endif

         ACC_TYPE score[MQ_GQA];
         #pragma unroll
         for (int h = 0; h < MQ_GQA; ++h) {
             const ACC_TYPE dot_partial = dot4[h].s0 + dot4[h].s1 + dot4[h].s2 + dot4[h].s3;
             ACC_TYPE s = sub_group_reduce_add(dot_partial) * scale;
+            if (logit_softcap > 0.0f) {
+                s = logit_softcap * tanh(s / logit_softcap);
+            }
             if (mask_base[h] != NULL) {
                 const global MASK_DATA_TYPE * mask_ptr = (const global MASK_DATA_TYPE *) mask_base[h];
                 s += slope[h] * (ACC_TYPE) mask_ptr[k_idx];
             }
-            if (logit_softcap > 0.0f) {
-                s = logit_softcap * tanh(s / logit_softcap);
-            }
             score[h] = s;
         }

@@ -1810,7 +1842,7 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split(
         barrier(CLK_LOCAL_MEM_FENCE);

         if (sgid == 0) {
-            const int head_idx = head_kv_idx * MQ_GQA + h;
+            const int head_idx = FA_MQS_HEAD_IDX(h);

             // fold per-subgroup (m, l) into split-level (m_c, l_c)
             ACC_TYPE m_c = sg_m[h][0];
@@ -1848,6 +1880,9 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split(
     }
 }

+#undef FA_MQS_HEAD_IDX
+
+#ifndef FA_MQ_SPLIT_ONLY
 // Cluster-parallel variant of _q1_vec_mq_split
 //
 // Tthe baseline keeps one 256B K row in flight per subgroup (32 lanes cooperate
@@ -1936,8 +1971,12 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split_c8(
     const int split_idx        = split_q_idx % n_splits;
     const int q_idx            = split_q_idx / n_splits;

-    const int batch_idx   = kvhead_batch_idx / n_head_kv;
-    const int head_kv_idx = kvhead_batch_idx % n_head_kv;
+    const int hgroups     = n_head_kv * FA_HEAD_SUB;
+    const int batch_idx   = kvhead_batch_idx / hgroups;
+    const int hg          = kvhead_batch_idx % hgroups;
+    const int head_kv_idx = hg / FA_HEAD_SUB;
+    const int head_sub    = hg % FA_HEAD_SUB;
+#define FA_HEAD_IDX(h) (head_kv_idx * (MQ_GQA * FA_HEAD_SUB) + head_sub * MQ_GQA + (h))

     const int kv_start = split_idx * kv_per_split;
     const int kv_end   = min(kv_start + kv_per_split, n_kv);
@@ -1948,7 +1987,7 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split_c8(
         if (tid == 0) {
             #pragma unroll
             for (int h = 0; h < MQ_GQA; ++h) {
-                const int head_idx = head_kv_idx * MQ_GQA + h;
+                const int head_idx = FA_HEAD_IDX(h);
                 const ulong rec_idx = ((((ulong) batch_idx * n_head + head_idx) * n_q + q_idx)
                                        * n_splits + split_idx);
                 global float * rec = partial_void + rec_idx * record_stride;
@@ -1968,7 +2007,7 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split_c8(
     for (int i = tid; i < MQ_GQA * DK_VEC; i += MQ_SPLIT_WG_SIZE) {
         const int h        = i / DK_VEC;
         const int k        = i % DK_VEC;
-        const int head_idx = head_kv_idx * MQ_GQA + h;
+        const int head_idx = FA_HEAD_IDX(h);
         const ulong q_row_offset = batch_idx * q_nb3 + head_idx * q_nb2 + (ulong) q_idx * q_nb1;
         const global Q_DATA_TYPE4 * q_ptr = (const global Q_DATA_TYPE4 *) (q_base + q_row_offset);
         q_shared[h * DK_VEC + k] = CONVERT_Q_ACC4(q_ptr[k]);
@@ -1978,9 +2017,17 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split_c8(
     float slope[MQ_GQA];
     #pragma unroll
     for (int h = 0; h < MQ_GQA; ++h) {
-        slope[h] = get_alibi_slope(max_bias, head_kv_idx * MQ_GQA + h, n_head_log2, m0, m1);
+        slope[h] = get_alibi_slope(max_bias, FA_HEAD_IDX(h), n_head_log2, m0, m1);
     }

+#ifdef FA_CL_MASK_BCAST
+    const global char * mask_base_b = NULL;
+    if (mask_void != NULL) {
+        mask_base_b = (const global char *) mask_void + mask_offset +
+                      (batch_idx % mask_ne3) * mask_nb3 + (ulong) q_idx * mask_nb1;
+    }
+    const int mask_bcast = mask_base_b != NULL && mask_ne2 == 1;
+#else
     const global char * mask_base[MQ_GQA];
     if (mask_void != NULL) {
         const int mask_batch_idx = batch_idx % mask_ne3;
@@ -1989,7 +2036,7 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split_c8(
                                           (ulong) q_idx * mask_nb1;
         #pragma unroll
         for (int h = 0; h < MQ_GQA; ++h) {
-            const int head_idx      = head_kv_idx * MQ_GQA + h;
+            const int head_idx      = FA_HEAD_IDX(h);
             const int mask_head_idx = head_idx % mask_ne2;
             mask_base[h] = mask_base_b + mask_head_idx * mask_nb2;
         }
@@ -1997,6 +2044,7 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split_c8(
         #pragma unroll
         for (int h = 0; h < MQ_GQA; ++h) mask_base[h] = NULL;
     }
+#endif

     // Per-CLUSTER online-softmax state (uniform across the cluster's lanes);
     // o_acc holds this lane's DV slice {lic + FA_CL_C*i}.
@@ -2031,6 +2079,73 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split_c8(
         const global KV_DATA_TYPE4 * k_ptr = (const global KV_DATA_TYPE4 *) (k_base + kv_row_base + (ulong) k_safe * k_nb1);
         const global KV_DATA_TYPE4 * v_ptr = (const global KV_DATA_TYPE4 *) (v_base + v_row_base  + (ulong) k_safe * v_nb1);

+#if defined(FA_CL_MHRED) && MQ_GQA == 4 && FA_CL_C == 16 && FA_CL_DK == 1 && FA_CL_DV == 1
+        ACC_TYPE mask_val = 0.0f;
+        if (mask_bcast) {
+            mask_val = (ACC_TYPE) ((const global MASK_DATA_TYPE *) mask_base_b)[k_safe];
+        }
+        const ACC_TYPE4 k_vec_1 = CONVERT_KV_ACC4(k_ptr[lic]);
+        const ACC_TYPE4 v_vec_1 = CONVERT_KV_ACC4(v_ptr[lic]);
+
+        // Reduce four heads with eight shuffles and keep each head's summation order.
+        const int mh_b0 = lic & 1;
+        const int mh_b1 = lic & 2;
+
+        ACC_TYPE mh_p[MQ_GQA];
+        #pragma unroll
+        for (int h = 0; h < MQ_GQA; ++h) {
+            const ACC_TYPE4 d4 = mad(q_shared[h * DK_VEC + lic], k_vec_1, (ACC_TYPE4)(0.0f));
+            mh_p[h] = d4.s0 + d4.s1 + d4.s2 + d4.s3;
+        }
+
+        ACC_TYPE mh_r2[2];
+        #pragma unroll
+        for (int j = 0; j < 2; ++j) {
+            const ACC_TYPE keep = mh_b0 ? mh_p[j + 2] : mh_p[j];
+            const ACC_TYPE send = mh_b0 ? mh_p[j]     : mh_p[j + 2];
+            mh_r2[j] = keep + sub_group_shuffle_xor(send, 1);
+        }
+        ACC_TYPE mh_r1 = (mh_b1 ? mh_r2[1] : mh_r2[0]) +
+                         sub_group_shuffle_xor(mh_b1 ? mh_r2[0] : mh_r2[1], 2);
+        mh_r1 += sub_group_shuffle_xor(mh_r1, 4);
+        mh_r1 += sub_group_shuffle_xor(mh_r1, 8);
+
+        ACC_TYPE mh_e2[2];
+        {
+            const ACC_TYPE other = sub_group_shuffle_xor(mh_r1, 2);
+            mh_e2[0] = mh_b1 ? other  : mh_r1;
+            mh_e2[1] = mh_b1 ? mh_r1  : other;
+        }
+        ACC_TYPE mh_s[MQ_GQA];
+        #pragma unroll
+        for (int j = 0; j < 2; ++j) {
+            const ACC_TYPE other = sub_group_shuffle_xor(mh_e2[j], 1);
+            mh_s[j]     = mh_b0 ? other     : mh_e2[j];
+            mh_s[j + 2] = mh_b0 ? mh_e2[j]  : other;
+        }
+
+        #pragma unroll
+        for (int h = 0; h < MQ_GQA; ++h) {
+            ACC_TYPE s = mh_s[h] * scale;
+            if (logit_softcap > 0.0f) {
+                s = logit_softcap * tanh(s / logit_softcap);
+            }
+            if (mask_bcast) {
+                s += slope[h] * mask_val;
+            } else if (mask_base_b != NULL) {
+                const int mask_head_idx = FA_HEAD_IDX(h) % mask_ne2;
+                const global MASK_DATA_TYPE * mask_ptr = (const global MASK_DATA_TYPE *) (mask_base_b + mask_head_idx * mask_nb2);
+                s += slope[h] * (ACC_TYPE) mask_ptr[k_safe];
+            }
+            const ACC_TYPE sc    = valid ? s : FA_M_INIT;
+            const ACC_TYPE m_new = max(m_i[h], sc);
+            const ACC_TYPE sp    = native_exp(m_i[h] - m_new);
+            const ACC_TYPE p     = native_exp(sc - m_new);
+            l_i[h] = l_i[h] * sp + p;
+            m_i[h] = m_new;
+            o_acc[h][0] = mad(p, v_vec_1, o_acc[h][0] * sp);
+        }
+#else
         // Dot: this lane covers DK elements {lic + FA_CL_C*i} of the cluster's row.
         ACC_TYPE4 dot4[MQ_GQA];
         #pragma unroll
@@ -2055,13 +2170,13 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split_c8(
                 s += sub_group_shuffle_xor(s, step);
             }
             s *= scale;
+            if (logit_softcap > 0.0f) {
+                s = logit_softcap * tanh(s / logit_softcap);
+            }
             if (mask_base[h] != NULL) {
                 const global MASK_DATA_TYPE * mask_ptr = (const global MASK_DATA_TYPE *) mask_base[h];
                 s += slope[h] * (ACC_TYPE) mask_ptr[k_safe];
             }
-            if (logit_softcap > 0.0f) {
-                s = logit_softcap * tanh(s / logit_softcap);
-            }
             score[h] = valid ? s : FA_M_INIT;
         }

@@ -2087,6 +2202,7 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split_c8(
                 o_acc[h][i] = mad(p_h[h], v_vec, o_acc[h][i] * sp_h[h]);
             }
         }
+#endif
     }

     // Merge stage 1: fold the FA_CL_NCL cluster partials inside the subgroup.
@@ -2148,7 +2264,7 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split_c8(
         barrier(CLK_LOCAL_MEM_FENCE);

         if (sgid == 0) {
-            const int head_idx = head_kv_idx * MQ_GQA + h;
+            const int head_idx = FA_HEAD_IDX(h);

             ACC_TYPE m_c = sg_m[h][0];
             #pragma unroll
@@ -2184,6 +2300,8 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split_c8(
     }
 }

+#undef FA_HEAD_IDX
+
 #endif  // DK_VEC/DV_VEC divisible by FA_CL_C
 #endif  // HAS_SUBGROUP_SHUFFLE (q1_vec_mq_split_c8)

@@ -2419,9 +2537,11 @@ __kernel void flash_attn_f32_f16_q1_vec_mq_split_k_img(
         barrier(CLK_LOCAL_MEM_FENCE);
     }
 }
+#endif  // !FA_MQ_SPLIT_ONLY
 #endif  // !FA_DECODE_ONLY

 #ifndef FA_MQ_ONLY  // q1_split + merge excluded from the MQ-only (g8) program
+#ifndef FA_VEC_ONLY
 __kernel void flash_attn_f32_f16_q1_split(
     const global void * q_void, ulong q_offset,
     const global void * k_void, ulong k_offset,
@@ -2578,6 +2698,8 @@ __kernel void flash_attn_f32_f16_q1_split(
     }
 }

+#endif  // !FA_VEC_ONLY
+
 // FD Pass 2: merge per-split partials into final O
 // empty splits drop via exp(-INF)=0.
 __kernel void flash_attn_f32_merge(