From 307a87ae323bf20f9b184df2e0cf793e23e044a3 Mon Sep 17 00:00:00 2001
From: Mike Lothian <mike@fireburn.co.uk>
Date: Sun, 27 Sep 2026 13:29:47 +0100
Subject: [PATCH 16/18] vulkan : load 16 values per thread for q4_K/q5_K/q6_K
 in mul_mm, use it for q6_K on RDNA4

Assisted-by: Claude Code (Claude Opus 5.5)
---
 ggml/src/ggml-vulkan/ggml-vulkan.cpp          |  2 +-
 .../ggml-vulkan/vulkan-shaders/mul_mm.comp    | 47 ++++++++-----
 .../vulkan-shaders/mul_mm_funcs.glsl          | 67 +++++++++++++++++++
 3 files changed, 97 insertions(+), 19 deletions(-)

diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
index 636848e05..438116249 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp
+++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
@@ -2493,7 +2493,7 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
             cm1_create_mmq({GGML_TYPE_Q3_K,   GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int_k, "matmul_q3_k_q8_1",   matmul_q3_k_q8_1_cm1_len,   matmul_q3_k_q8_1_cm1_data,   sizeof(vk_mat_mat_push_constants), 3);
             if (!rdna4) { cm1_create_mmq({GGML_TYPE_Q4_K, GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int,   "matmul_q4_k_q8_1",   matmul_q4_k_q8_1_cm1_len,   matmul_q4_k_q8_1_cm1_data,   sizeof(vk_mat_mat_push_constants), 3); }
             if (!rdna4) { cm1_create_mmq({GGML_TYPE_Q5_K, GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int,   "matmul_q5_k_q8_1",   matmul_q5_k_q8_1_cm1_len,   matmul_q5_k_q8_1_cm1_data,   sizeof(vk_mat_mat_push_constants), 3); }
-            cm1_create_mmq({GGML_TYPE_Q6_K,   GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int_k, "matmul_q6_k_q8_1",   matmul_q6_k_q8_1_cm1_len,   matmul_q6_k_q8_1_cm1_data,   sizeof(vk_mat_mat_push_constants), 3);
+            if (!rdna4) { cm1_create_mmq({GGML_TYPE_Q6_K, GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int_k, "matmul_q6_k_q8_1",   matmul_q6_k_q8_1_cm1_len,   matmul_q6_k_q8_1_cm1_data,   sizeof(vk_mat_mat_push_constants), 3); }
             if (!rdna4) { cm1_create_mmq({GGML_TYPE_NVFP4, GGML_TYPE_Q8_1, false, false}, tc_mmq_cm1_int_k, "matmul_nvfp4_q8_1",  matmul_nvfp4_q8_1_cm1_len,  matmul_nvfp4_q8_1_cm1_data,  sizeof(vk_mat_mat_push_constants), 3); }
         }
 
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp
index 11098ee7b..31a57b18c 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm.comp
@@ -52,24 +52,6 @@ layout (constant_id = 11) const uint ALIGNED = 0;
 
 #ifdef MULMAT_QUANT
 
-uint mm_load_vec_a() {
-    switch (MmTypeA) {
-    case GGML_TYPE_Q1_0:
-    case GGML_TYPE_Q4_0:
-    case GGML_TYPE_Q4_1:
-    case GGML_TYPE_Q5_1:
-        return 8u;
-    case GGML_TYPE_Q2_0:
-    case GGML_TYPE_Q5_0:
-    case GGML_TYPE_Q8_0:
-    case GGML_TYPE_Q2_K:
-    case GGML_TYPE_Q4_K:
-    case GGML_TYPE_Q5_K:
-        return 4u;
-    default:
-        return 2u;
-    }
-}
 #endif
 
 #if !defined(TO_FLOAT_TYPE)
@@ -177,6 +159,35 @@ layout (constant_id = 3) const uint BK = 16;  // Assumed to be 32 if working wit
 #define BK_STEP 2
 #endif
 
+#ifdef MULMAT_QUANT
+uint mm_load_vec_a() {
+    switch (MmTypeA) {
+    case GGML_TYPE_Q1_0:
+    case GGML_TYPE_Q4_0:
+    case GGML_TYPE_Q4_1:
+    case GGML_TYPE_Q5_1:
+        return 8u;
+    case GGML_TYPE_Q2_0:
+    case GGML_TYPE_Q5_0:
+    case GGML_TYPE_Q8_0:
+    case GGML_TYPE_Q2_K:
+        return 4u;
+    case GGML_TYPE_Q4_K:
+    case GGML_TYPE_Q5_K:
+    case GGML_TYPE_Q6_K:
+#ifndef MM_KQ16_DISABLE
+        // 16 values of one sub-block per thread: one scale decode instead of four
+        if (BLOCK_SIZE * 16 / BK <= BM && BM % (BLOCK_SIZE * 16 / BK) == 0) {
+            return 16u;
+        }
+#endif
+        return MmTypeA == GGML_TYPE_Q6_K ? 2u : 4u;
+    default:
+        return 2u;
+    }
+}
+#endif
+
 #ifdef COOPMAT
 #ifdef MULMAT_QUANT
 layout(constant_id = 13) const uint SHMEM_STRIDE_PAD = 4;
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl
index 588fb7354..5500f7efb 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mm_funcs.glsl
@@ -501,6 +501,52 @@ void load_a_to_shmem(const uint pos_a, const uint row, const uint col, const uin
         store_a(col, k_pair, FLOAT_TYPEV2(dl * (qs.x - hm.x),
                                         dl * (qs.y - hm.y)));
 
+    } else if ((MmTypeA == GGML_TYPE_Q4_K || MmTypeA == GGML_TYPE_Q5_K) && mm_load_vec_a() == 16) {
+        const uint idx = pos_a + col * p.stride_a / 16 + row;
+        const uint k_pair = row * 8;
+
+        const uint ib = idx / 16;                  // 16 values per idx
+        const uint k0 = (idx % 16) * 16;           // 0,16..240
+        const uint n = k0 / 64;                    // 0..3
+        const uint b = (k0 % 64) / 32;             // 0,1
+        const uint is = 2 * n + b;                 // 0..7
+        const uint qsi = (n * 32 + k0 % 32) / 4;   // first of 4 qs dwords
+
+        vec2 loadd;
+        uvec3 scales;
+        uvec4 qs;
+        uvec4 qh = uvec4(0);
+        if (MmTypeA == GGML_TYPE_Q4_K) {
+            loadd = vec2(a_q4_k.data[ib].dm);
+            scales = uvec3(a_q4_k_p32.data[ib].scales[0], a_q4_k_p32.data[ib].scales[1], a_q4_k_p32.data[ib].scales[2]);
+            qs = uvec4(a_q4_k_p32.data[ib].qs[qsi], a_q4_k_p32.data[ib].qs[qsi + 1],
+                       a_q4_k_p32.data[ib].qs[qsi + 2], a_q4_k_p32.data[ib].qs[qsi + 3]);
+        } else {
+            loadd = vec2(a_q5_k.data[ib].dm);
+            scales = uvec3(a_q5_k_p32.data[ib].scales[0], a_q5_k_p32.data[ib].scales[1], a_q5_k_p32.data[ib].scales[2]);
+            qs = uvec4(a_q5_k_p32.data[ib].qs[qsi], a_q5_k_p32.data[ib].qs[qsi + 1],
+                       a_q5_k_p32.data[ib].qs[qsi + 2], a_q5_k_p32.data[ib].qs[qsi + 3]);
+            const uint qhi = (k0 % 32) / 4;
+            qh = (uvec4(a_q5_k_p32.data[ib].qh[qhi], a_q5_k_p32.data[ib].qh[qhi + 1],
+                        a_q5_k_p32.data[ib].qh[qhi + 2], a_q5_k_p32.data[ib].qh[qhi + 3]) >> (k0 / 32)) & 0x01010101u;
+            qh <<= 4;
+        }
+        const uint scalesoffs = (is & 3) * 8;
+        const uint scidx0 = (is < 4) ? 0 : 2;
+        const uint scidxshift1 = (is < 4) ? scalesoffs : scalesoffs + 2;
+        const uint mbidx0 = (is < 4) ? 1 : 2;
+        const uint mbidxshift0 = (is < 4) ? scalesoffs : scalesoffs + 4;
+        const uint mbidxshift1 = (is < 4) ? scalesoffs : scalesoffs + 2;
+        const uint sc    = ((scales[scidx0] >> scalesoffs) & 0xF) | ((scales[0] >> scidxshift1) & 0x30);
+        const uint mbyte = ((scales[mbidx0] >> mbidxshift0) & 0xF) | ((scales[1] >> mbidxshift1) & 0x30);
+        const float d = loadd.x * float(sc);
+        const float m = -loadd.y * float(mbyte);
+
+        [[unroll]] for (uint w = 0; w < 4; ++w) {
+            const vec4 q = vec4(unpack8(((qs[w] >> (b * 4)) & 0x0F0F0F0Fu) | qh[w]));
+            store_a(col, k_pair + 2 * w,     FLOAT_TYPEV2(fma(d, q.x, m), fma(d, q.y, m)));
+            store_a(col, k_pair + 2 * w + 1, FLOAT_TYPEV2(fma(d, q.z, m), fma(d, q.w, m)));
+        }
     } else if (MmTypeA == GGML_TYPE_Q4_K) {
         const uint idx = pos_a + col * p.stride_a / mm_load_vec_a() + row;
         const uint k_pair = row * mm_load_vec_a() / 2;
@@ -576,6 +622,27 @@ void load_a_to_shmem(const uint pos_a, const uint row, const uint col, const uin
 
         store_a(col, k_pair, FLOAT_TYPEV2(fma(d, q.x, m), fma(d, q.y, m)));
         store_a(col, k_pair + 1, FLOAT_TYPEV2(fma(d, q.z, m), fma(d, q.w, m)));
+    } else if (MmTypeA == GGML_TYPE_Q6_K && mm_load_vec_a() == 16) {
+        const uint idx = pos_a + col * p.stride_a / 16 + row;
+        const uint k_pair = row * 8;
+
+        const uint ib = idx / 16;                   // 16 values per idx
+        const uint iqs = (idx % 16) * 8;            // k0 / 2
+        const uint n = iqs / 64;                    // 0,1
+        const uint b = ((iqs % 64) / 32) * 4;       // 0,4
+        const uint qhshift = ((iqs % 64) / 16) * 2; // 0,2,4,6
+        const uint is = 8 * n + qhshift + (iqs % 16) / 8;
+        const uint qsi = n * 32 + (iqs % 32);
+        const uint qhi = n * 16 + (iqs % 16);
+
+        const float dscale = float(a_q6_k.data[ib].d) * float(a_q6_k.data[ib].scales[is]);
+
+        [[unroll]] for (uint w = 0; w < 8; ++w) {
+            const uint ql = (uint(a_q6_k_p16.data[ib].ql[qsi + w]) >> b) & 0x0F0F;
+            const uint qh = (uint(a_q6_k_p16.data[ib].qh[qhi + w]) >> qhshift) & 0x0303;
+            const vec2 q = (vec2(unpack8(ql | (qh << 4)).xy) - 32) * dscale;
+            store_a(col, k_pair + w, FLOAT_TYPEV2(q.x, q.y));
+        }
     } else if (MmTypeA == GGML_TYPE_Q6_K) {
         const uint idx = pos_a + col * p.stride_a / mm_load_vec_a() + row;
         const uint k_pair = row * mm_load_vec_a() / 2;
-- 
2.55.0

