From b0c54beddc877111fe3513d6a630baf066045a18 Mon Sep 17 00:00:00 2001
From: Mike Lothian <mike@fireburn.co.uk>
Date: Sat, 26 Sep 2026 15:38:21 +0100
Subject: [PATCH 09/18] hip : share the FA vec K/V pass between the Q heads of
 a GQA group on RDNA

Assisted-by: Claude Code (Claude Opus 5.5)
---
 ggml/src/ggml-cuda/fattn-vec.cuh | 128 +++++++++++++++++++++----------
 tests/test-backend-ops.cpp       |  16 ++++
 2 files changed, 105 insertions(+), 39 deletions(-)

diff --git a/ggml/src/ggml-cuda/fattn-vec.cuh b/ggml/src/ggml-cuda/fattn-vec.cuh
index 57a285565..67e454fd4 100644
--- a/ggml/src/ggml-cuda/fattn-vec.cuh
+++ b/ggml/src/ggml-cuda/fattn-vec.cuh
@@ -16,7 +16,9 @@ static constexpr __device__ int ggml_cuda_fattn_vec_get_nthreads_device() {
 #pragma clang diagnostic push
 #pragma clang diagnostic ignored "-Wpass-failed"
 #endif // __clang__
-template<int D, int ncols, ggml_type type_K, ggml_type type_V, bool use_logit_softcap> // D == head size
+// ncols1: Q columns (tokens) per block, ncols2: Q heads of the same GQA group per block.
+// With ncols2 > 1 the K/V data is loaded once for all ncols2 heads.
+template<int D, int ncols1, int ncols2, ggml_type type_K, ggml_type type_V, bool use_logit_softcap> // D == head size
 __launch_bounds__(ggml_cuda_fattn_vec_get_nthreads_device(), 1)
 static __global__ void flash_attn_ext_vec(
         const char * Q_ptr,
@@ -68,12 +70,15 @@ static __global__ void flash_attn_ext_vec(
 
     //In this kernel Q, K, V are matrices while i, j, k are matrix indices.
 
+    constexpr int ncols = ncols1*ncols2;
+
     constexpr int cpy_nb = ggml_cuda_get_max_cpy_bytes();
     constexpr int cpy_ne = cpy_nb / 4;
 
 #ifdef GGML_USE_HIP
 #ifdef RDNA
-    constexpr int nthreads_KQ_q = 2;
+    // With several heads per block Q must not be held whole by 2 threads per column (VGPR budget).
+    constexpr int nthreads_KQ_q = ncols2 > 1 ? 8 : 2;
 #else
     constexpr int nthreads_KQ_q = 4;
 #endif // RDNA
@@ -101,18 +106,19 @@ static __global__ void flash_attn_ext_vec(
     constexpr dequantize_V_t dequantize_V = get_dequantize_V<type_V, float, V_rows_per_thread>();
 #endif // V_DOT2_F32_F16_AVAILABLE
 
-    const int ic0 = blockIdx.x * ncols; // Index of the Q/QKV column to work on.
+    const int ic0 = blockIdx.x * ncols1; // Index of the first Q/QKV column to work on.
 
-    const int sequence = blockIdx.z / ne02;
-    const int head = blockIdx.z - sequence*ne02;
+    const int sequence = blockIdx.z / (ne02/ncols2);
+    const int head0 = blockIdx.z*ncols2 - sequence*ne02; // First Q head of this block.
     const int gqa_ratio = ne02 / ne12; // With grouped query attention there are > 1 Q matrices per K, V matrix.
-    Q += nb03*sequence + nb02* head              + nb01*ic0;
-    K += nb13*sequence + nb12*(head / gqa_ratio);
-    V += nb23*sequence + nb22*(head / gqa_ratio);
+    Q += nb03*sequence + nb02*head0              + nb01*ic0;
+    K += nb13*sequence + nb12*(head0 / gqa_ratio);
+    V += nb23*sequence + nb22*(head0 / gqa_ratio);
 
     const half * maskh  = (const half  *) (mask + nb33*(sequence % ne33) + nb31*ic0);
 
-    const float slope = get_alibi_slope(max_bias, head, n_head_log2, m0, m1);
+    // ALiBi is excluded on the host when ncols2 > 1.
+    const float slope = ncols2 == 1 ? get_alibi_slope(max_bias, head0, n_head_log2, m0, m1) : 1.0f;
 
     static_assert(D % (2*WARP_SIZE) == 0, "D not divisible by 2*WARP_SIZE == 64.");
     constexpr int nwarps = nthreads / WARP_SIZE;
@@ -161,7 +167,7 @@ static __global__ void flash_attn_ext_vec(
             float2 * tmp_q_ds  = (float2 *) (tmp_q_i32 + D/sizeof(int));
 
             // Set memory to zero if out of bounds:
-            if (ncols > 1 && ic0 + j >= int(ne01.z)) {
+            if (ncols1 > 1 && ic0 + j/ncols2 >= int(ne01.z)) {
 #pragma unroll
                 for (int i0 = 0; i0 < int(D/sizeof(int)); i0 += WARP_SIZE) {
                     const int i = i0 + threadIdx.x;
@@ -174,7 +180,7 @@ static __global__ void flash_attn_ext_vec(
                     tmp_q_ds[threadIdx.x] = make_float2(0.0f, 0.0f);
                 }
             } else {
-                const float * Q_f = (const float *) (Q + j*nb01);
+                const float * Q_f = (const float *) (Q + (j/ncols2)*nb01 + (j%ncols2)*nb02);
                 constexpr int nthreads_quantize = D/sizeof(int) < WARP_SIZE ? D/sizeof(int) : WARP_SIZE;
 #pragma unroll
                 for (int i0 = 0; i0 < int(D/sizeof(int)); i0 += nthreads_quantize) {
@@ -206,13 +212,13 @@ static __global__ void flash_attn_ext_vec(
         const half2 scale_h2 = make_half2(scale, scale);
 #pragma unroll
         for (int j = 0; j < ncols; ++j) {
-            const float2 * Q_j = (const float2 *) (Q + j*nb01);
+            const float2 * Q_j = (const float2 *) (Q + (j/ncols2)*nb01 + (j%ncols2)*nb02);
 #pragma unroll
             for (int i0 = 0; i0 < D/2; i0 += nthreads_KQ*cpy_ne) {
                 const int i = i0 + (nthreads_KQ == WARP_SIZE ? threadIdx.x : threadIdx.x % nthreads_KQ)*cpy_ne;
 
                 __align__(16) float2 tmp[cpy_ne] = {{0.0f, 0.0f}};
-                if (ncols == 1 || ic0 + j < int(ne01.z)) {
+                if (ncols1 == 1 || ic0 + j/ncols2 < int(ne01.z)) {
                     ggml_cuda_memcpy_1<cpy_nb>(tmp,            &Q_j[i]);
                     ggml_cuda_memcpy_1<cpy_nb>(tmp + cpy_ne/2, &Q_j[i + cpy_ne/2]);
                 }
@@ -229,11 +235,11 @@ static __global__ void flash_attn_ext_vec(
 #else
 #pragma unroll
         for (int j = 0; j < ncols; ++j) {
-            const float2 * Q_j = (const float2 *) (Q + j*nb01);
+            const float2 * Q_j = (const float2 *) (Q + (j/ncols2)*nb01 + (j%ncols2)*nb02);
 #pragma unroll
             for (int i0 = 0; i0 < D/2; i0 += nthreads_KQ*cpy_ne) {
                 const int i = i0 + (nthreads_KQ == WARP_SIZE ? threadIdx.x : threadIdx.x % nthreads_KQ)*cpy_ne;
-                if (ncols == 1 || ic0 + j < int(ne01.z)) {
+                if (ncols1 == 1 || ic0 + j/ncols2 < int(ne01.z)) {
                     ggml_cuda_memcpy_1<cpy_nb>(&Q_reg[j][i0/nthreads_KQ],            &Q_j[i]);
                     ggml_cuda_memcpy_1<cpy_nb>(&Q_reg[j][i0/nthreads_KQ + cpy_ne/2], &Q_j[i + cpy_ne/2]);
                 }
@@ -277,8 +283,8 @@ static __global__ void flash_attn_ext_vec(
                     sum = logit_softcap*tanhf(sum);
                 }
 
-                if (mask && (ncols == 1 || ic0 + j < int(ne01.z))) {
-                    sum += slope*__half2float(maskh[j*ne11 + i_KQ]);
+                if (mask && (ncols1 == 1 || ic0 + j/ncols2 < int(ne01.z))) {
+                    sum += slope*__half2float(maskh[(j/ncols2)*ne11 + i_KQ]);
                 }
 
                 KQ_max_new[j] = fmaxf(KQ_max_new[j], sum + FATTN_KQ_MAX_OFFSET);
@@ -377,8 +383,6 @@ static __global__ void flash_attn_ext_vec(
     }
 
     if (sinks && blockIdx.y == 0) {
-        const float sink = ((const float *) sinks)[head];
-
 #pragma unroll
         for (int j0 = 0; j0 < ncols; j0 += nwarps) {
             const int j = j0 + threadIdx.y;
@@ -387,6 +391,8 @@ static __global__ void flash_attn_ext_vec(
                 break;
             }
 
+            const float sink = ((const float *) sinks)[head0 + j%ncols2];
+
             const float kqmax_new_j = fmaxf(sink, KQ_max[j]);
             const float KQ_max_scale = expf(KQ_max[j] - kqmax_new_j);
             KQ_max[j] = kqmax_new_j;
@@ -431,7 +437,7 @@ static __global__ void flash_attn_ext_vec(
 
 #pragma unroll
     for (int j_VKQ = 0; j_VKQ < ncols; ++j_VKQ) {
-        if (ncols > 1 && ic0 + j_VKQ >= int(ne01.z)) {
+        if (ncols1 > 1 && ic0 + j_VKQ/ncols2 >= int(ne01.z)) {
             break;
         }
 
@@ -498,7 +504,7 @@ static __global__ void flash_attn_ext_vec(
                 if (gridDim.y == 1) {
                     dst_val /= KQ_sum[j_VKQ];
                 }
-                dst[(((sequence*int(ne01.z) + ic0 + j_VKQ)*ne02 + head)*gridDim.y + blockIdx.y)*D + i0 + tid] = dst_val;
+                dst[(((sequence*int(ne01.z) + ic0 + j_VKQ/ncols2)*ne02 + head0 + j_VKQ%ncols2)*gridDim.y + blockIdx.y)*D + i0 + tid] = dst_val;
             }
         }
 
@@ -508,8 +514,8 @@ static __global__ void flash_attn_ext_vec(
 
     }
 
-    if (gridDim.y != 1 && tid < ncols && (ncols == 1 || ic0 + tid < int(ne01.z))) {
-        dst_meta[((sequence*int(ne01.z) + ic0 + tid)*ne02 + head)*gridDim.y + blockIdx.y] = make_float2(KQ_max[tid], KQ_sum[tid]);
+    if (gridDim.y != 1 && tid < ncols && (ncols1 == 1 || ic0 + tid/ncols2 < int(ne01.z))) {
+        dst_meta[((sequence*int(ne01.z) + ic0 + tid/ncols2)*ne02 + head0 + tid%ncols2)*gridDim.y + blockIdx.y] = make_float2(KQ_max[tid], KQ_sum[tid]);
     }
 #else
     GGML_UNUSED_VARS(Q_ptr, K_ptr, V_ptr, mask_ptr, sinks_ptr, KV_max_ptr, dst_ptr, dst_meta_ptr, scale,
@@ -528,17 +534,52 @@ static __global__ void flash_attn_ext_vec(
 #pragma clang diagnostic pop
 #endif // __clang__
 
-template <int D, int cols_per_block, ggml_type type_K, ggml_type type_V, bool use_logit_softcap>
+template <int D, int cols_per_block, int ncols2, ggml_type type_K, ggml_type type_V, bool use_logit_softcap>
 void ggml_cuda_flash_attn_ext_vec_case_impl(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
     const int cc = ggml_cuda_info().devices[ggml_cuda_get_device()].cc;
 
     const int nthreads = ggml_cuda_fattn_vec_get_nthreads_host(cc);
     const int nwarps   = nthreads / WARP_SIZE;
-    fattn_kernel_t fattn_kernel = flash_attn_ext_vec<D, cols_per_block, type_K, type_V, use_logit_softcap>;
+    fattn_kernel_t fattn_kernel = flash_attn_ext_vec<D, cols_per_block, ncols2, type_K, type_V, use_logit_softcap>;
     const bool need_f16_K = type_K == GGML_TYPE_F16;
     const bool need_f16_V = type_V == GGML_TYPE_F16;
     constexpr size_t nbytes_shared = 0;
-    launch_fattn<D, cols_per_block, 1>(ctx, dst, fattn_kernel, nwarps, nbytes_shared, D, need_f16_K, need_f16_V, false, false);
+    launch_fattn<D, cols_per_block, ncols2>(ctx, dst, fattn_kernel, nwarps, nbytes_shared, D, need_f16_K, need_f16_V, false, false);
+}
+
+template <int D, int cols_per_block, int ncols2, ggml_type type_K, ggml_type type_V>
+static void ggml_cuda_flash_attn_ext_vec_case_softcap(ggml_backend_cuda_context & ctx, ggml_tensor * dst, const float logit_softcap) {
+    if (logit_softcap == 0.0f) {
+        ggml_cuda_flash_attn_ext_vec_case_impl<D, cols_per_block, ncols2, type_K, type_V, false>(ctx, dst);
+    } else {
+        ggml_cuda_flash_attn_ext_vec_case_impl<D, cols_per_block, ncols2, type_K, type_V, true>(ctx, dst);
+    }
+}
+
+// Number of Q heads of a GQA group to process per block, 1 if the packing can not be used.
+static int ggml_cuda_fattn_vec_get_ncols2(const ggml_tensor * dst) {
+    const ggml_tensor * Q    = dst->src[0];
+    const ggml_tensor * K    = dst->src[1];
+    const ggml_tensor * V    = dst->src[2];
+    const ggml_tensor * mask = dst->src[3];
+
+    float max_bias = 0.0f;
+    memcpy(&max_bias, (const float *) dst->op_params + 1, sizeof(float));
+
+    const int gqa_ratio = Q->ne[2] / K->ne[2];
+    if (!mask || max_bias != 0.0f || Q->ne[1] != 1) {
+        return 1;
+    }
+    // While K/V fit in the last level cache, re-reading them per Q head is cheaper than the packed kernel.
+    if (ggml_nbytes(K) + ggml_nbytes(V) < 16*1024*1024) {
+        return 1;
+    }
+    for (const int ncols2 : {8, 6, 4, 2}) {
+        if (gqa_ratio % ncols2 == 0) {
+            return ncols2;
+        }
+    }
+    return 1;
 }
 
 template <int D, ggml_type type_K, ggml_type type_V>
@@ -551,24 +592,33 @@ void ggml_cuda_flash_attn_ext_vec_case(ggml_backend_cuda_context & ctx, ggml_ten
 
     if (Q->ne[1] == 1) {
         constexpr int cols_per_block = 1;
-        if (logit_softcap == 0.0f) {
-            constexpr bool use_logit_softcap = false;
-            ggml_cuda_flash_attn_ext_vec_case_impl<D, cols_per_block, type_K, type_V, use_logit_softcap>(ctx, dst);
-        } else {
-            constexpr bool use_logit_softcap = true;
-            ggml_cuda_flash_attn_ext_vec_case_impl<D, cols_per_block, type_K, type_V, use_logit_softcap>(ctx, dst);
+#ifdef GGML_USE_HIP
+        // Share the K/V pass between the Q heads of a GQA group (quantized K/V decode is memory bound).
+        const int cc = ggml_cuda_info().devices[ggml_cuda_get_device()].cc;
+        const int ncols2 = GGML_CUDA_CC_IS_RDNA(cc) ? ggml_cuda_fattn_vec_get_ncols2(dst) : 1;
+        if (ncols2 == 8) {
+            ggml_cuda_flash_attn_ext_vec_case_softcap<D, cols_per_block, 8, type_K, type_V>(ctx, dst, logit_softcap);
+            return;
+        }
+        if (ncols2 == 6) {
+            ggml_cuda_flash_attn_ext_vec_case_softcap<D, cols_per_block, 6, type_K, type_V>(ctx, dst, logit_softcap);
+            return;
         }
+        if (ncols2 == 4) {
+            ggml_cuda_flash_attn_ext_vec_case_softcap<D, cols_per_block, 4, type_K, type_V>(ctx, dst, logit_softcap);
+            return;
+        }
+        if (ncols2 == 2) {
+            ggml_cuda_flash_attn_ext_vec_case_softcap<D, cols_per_block, 2, type_K, type_V>(ctx, dst, logit_softcap);
+            return;
+        }
+#endif // GGML_USE_HIP
+        ggml_cuda_flash_attn_ext_vec_case_softcap<D, cols_per_block, 1, type_K, type_V>(ctx, dst, logit_softcap);
         return;
     }
 
     constexpr int cols_per_block = 2;
-    if (logit_softcap == 0.0f) {
-        constexpr bool use_logit_softcap = false;
-        ggml_cuda_flash_attn_ext_vec_case_impl<D, cols_per_block, type_K, type_V, use_logit_softcap>(ctx, dst);
-    } else {
-        constexpr bool use_logit_softcap = true;
-        ggml_cuda_flash_attn_ext_vec_case_impl<D, cols_per_block, type_K, type_V, use_logit_softcap>(ctx, dst);
-    }
+    ggml_cuda_flash_attn_ext_vec_case_softcap<D, cols_per_block, 1, type_K, type_V>(ctx, dst, logit_softcap);
 }
 
 #define DECL_FATTN_VEC_CASE(D, type_K, type_V)                              \
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index 32c3f3d2a..7bb3fca89 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -10970,6 +10970,17 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
         }
     }
 
+    // quantized K/V decode with GQA ratios that pack Q heads per block (e.g. 24 Q / 4 KV heads)
+    for (ggml_type type_KV : { GGML_TYPE_Q8_0, GGML_TYPE_Q4_0 }) {
+        for (int hs : { 128, 256, }) {
+            for (int gqa : { 2, 4, 6, 8, }) {
+                for (bool sinks : { false, true }) {
+                    test_cases.emplace_back(new test_flash_attn_ext(hs, hs, 4, {gqa, 1}, 32768, 1, true, sinks, 0, 0, GGML_PREC_F32, type_KV, type_KV));
+                }
+            }
+        }
+    }
+
     // asymmetric head_dim (hsk != hsv) with one or both sides not 64-aligned
     test_cases.emplace_back(new test_flash_attn_ext(72, 64, 4, {1, 1}, 256, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
     test_cases.emplace_back(new test_flash_attn_ext(64, 72, 4, {1, 1}, 256, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
@@ -11261,6 +11272,11 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
 static std::vector<std::unique_ptr<test_case>> make_test_cases_perf() {
     std::vector<std::unique_ptr<test_case>> test_cases;
 
+    // Qwen3.x 27B style decode: 24 Q heads / 4 KV heads, head size 256, q8_0 KV cache
+    for (int kv : { 1024, 4096, 8192, 16384, 65536, }) {
+        test_cases.emplace_back(new test_flash_attn_ext(256, 256, 4, {6, 1}, kv, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0));
+    }
+
     // SWIGLU at a 27B-class FFN width, fused [gate|up] vs split operands
     // note: same bytes either way, so a backend that indexes them differently shows it here
     for (ggml_type type : {GGML_TYPE_F16, GGML_TYPE_F32}) {
-- 
2.55.0

