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(