From 42c46aca04327a8c078320d7624c3a42484b2804 Mon Sep 17 00:00:00 2001
From: Mike Lothian <mike@fireburn.co.uk>
Date: Sun, 27 Sep 2026 07:27:30 +0100
Subject: [PATCH 18/18] vulkan : index iq4/mxfp4 codebooks as a constant array
 in cm1 mmq on RADV

Needs Mesa MR !44752, otherwise set GGML_VK_DISABLE_CONST_LUT=1.

Assisted-by: Claude Code (Claude Opus 5.5)
---
 ggml/src/ggml-vulkan/ggml-vulkan.cpp          | 10 +++--
 .../vulkan-shaders/mul_mmq_cm1.comp           | 17 +++++---
 .../vulkan-shaders/mul_mmq_cm1_funcs.glsl     | 40 +++++++++++--------
 3 files changed, 42 insertions(+), 25 deletions(-)

diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
index 438116249..8a89d320d 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp
+++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
@@ -1809,9 +1809,13 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
             return cm1_sg * (bm / std::min(cm1_sg, bm)) * (bn / 32);
         };
 
-        l_warptile_mmq_cm1_int = { cm1_bs(128, 128), 128, 128, 32, std::min(cm1_sg, 128u), 32, 2, itm, itn, itk, cm1_sg, (uint32_t)device->architecture };
-        m_warptile_mmq_cm1_int = { cm1_bs( 64,  64),  64,  64, 32, std::min(cm1_sg,  64u), 32, 2, itm, itn, itk, cm1_sg, (uint32_t)device->architecture };
-        s_warptile_mmq_cm1_int = { cm1_bs( 32,  32),  32,  32, 32, std::min(cm1_sg,  32u), 32, 2, itm, itn, itk, cm1_sg, (uint32_t)device->architecture };
+        // RADV lowers small constant byte tables to v_perm_b32 with Mesa MR !44752.
+        // Without it the lookup is a bcsel ladder, about 2x slower: set GGML_VK_DISABLE_CONST_LUT=1.
+        const uint32_t cm1_const_lut = device->driver_id == vk::DriverId::eMesaRadv &&
+                                       getenv("GGML_VK_DISABLE_CONST_LUT") == nullptr;
+        l_warptile_mmq_cm1_int = { cm1_bs(128, 128), 128, 128, 32, std::min(cm1_sg, 128u), 32, 2, itm, itn, itk, cm1_sg, (uint32_t)device->architecture, cm1_const_lut };
+        m_warptile_mmq_cm1_int = { cm1_bs( 64,  64),  64,  64, 32, std::min(cm1_sg,  64u), 32, 2, itm, itn, itk, cm1_sg, (uint32_t)device->architecture, cm1_const_lut };
+        s_warptile_mmq_cm1_int = { cm1_bs( 32,  32),  32,  32, 32, std::min(cm1_sg,  32u), 32, 2, itm, itn, itk, cm1_sg, (uint32_t)device->architecture, cm1_const_lut };
 
         l_warptile_mmq_cm1_int_k = { cm1_bs( 64, 128),  64, 128, 32, std::min(cm1_sg,  64u), 32, 2, itm, itn, itk, cm1_sg, (uint32_t)device->architecture };
         m_warptile_mmq_cm1_int_k = { cm1_bs( 64,  64),  64,  64, 32, std::min(cm1_sg,  64u), 32, 2, itm, itn, itk, cm1_sg, (uint32_t)device->architecture };
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp
index 9f0596ce7..562989af3 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1.comp
@@ -84,6 +84,7 @@ layout (constant_id = 8) const uint TN = 16;
 layout (constant_id = 9) const uint TK = 16;
 layout (constant_id = 10) const uint WARP = 32;
 layout (constant_id = 11) const uint DEVICE_ARCH = 0; // vk_device_architecture (ggml-vulkan.cpp)
+layout (constant_id = 12) const uint CONST_LUT = 0;   // look up iq4/mxfp4 codebooks as a constant array
 #define VK_ARCH_AMD_RDNA4 5u
 
 #define BK 32
@@ -157,15 +158,19 @@ ACC_TYPE cm1_accumulate(ACC_TYPE prev, int acc_e, float scale_a, float nbias_a,
 
 void main() {
 #if defined(DATA_A_IQ4_NL) || defined(DATA_A_IQ4_XS)
-    if (gl_LocalInvocationIndex < 16u) {
-        cm1_kvalues[gl_LocalInvocationIndex] = kvalues_iq4nl_const[gl_LocalInvocationIndex];
+    if (CONST_LUT == 0) {
+        if (gl_LocalInvocationIndex < 16u) {
+            cm1_kvalues[gl_LocalInvocationIndex] = kvalues_iq4nl_const[gl_LocalInvocationIndex];
+        }
+        barrier();
     }
-    barrier();
 #elif defined(DATA_A_MXFP4)
-    if (gl_LocalInvocationIndex < 16u) {
-        cm1_kvalues[gl_LocalInvocationIndex] = kvalues_mxfp4_const[gl_LocalInvocationIndex];
+    if (CONST_LUT == 0) {
+        if (gl_LocalInvocationIndex < 16u) {
+            cm1_kvalues[gl_LocalInvocationIndex] = kvalues_mxfp4_const[gl_LocalInvocationIndex];
+        }
+        barrier();
     }
-    barrier();
 #elif defined(DATA_A_NVFP4)
     if (gl_LocalInvocationIndex < 16u) {
         cm1_kvalues[gl_LocalInvocationIndex] = kvalues_mxfp4_const[gl_LocalInvocationIndex];
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl
index 1760c138e..0d92b98df 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mmq_cm1_funcs.glsl
@@ -1,3 +1,11 @@
+// CONST_LUT: index the codebook as a constant array (RADV lowers it to byte permutes), else read it from shared memory.
+#if defined(DATA_A_IQ4_NL) || defined(DATA_A_IQ4_XS)
+#define KV_LOOKUP(i) (CONST_LUT != 0 ? kvalues_iq4nl_const[i] : cm1_kvalues[i])
+#elif defined(DATA_A_MXFP4)
+#define KV_LOOKUP(i) (CONST_LUT != 0 ? kvalues_mxfp4_const[i] : cm1_kvalues[i])
+#else
+#define KV_LOOKUP(i) cm1_kvalues[i]
+#endif
 // Per-quant-type data structures and functions for the cm1 int8 coopmat path.
 // Each quant type defines:
 //   struct block_a_prefetch  — register data for one A-block per thread
@@ -162,11 +170,11 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
     const u8vec4 lo_idx = unpack8(blk.qs & 0x0F0F0F0F);
     const u8vec4 hi_idx = unpack8((blk.qs >> 4) & 0x0F0F0F0F);
     buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr    ] =
-        pack32(i8vec4(cm1_kvalues[lo_idx.x], cm1_kvalues[lo_idx.y],
-                      cm1_kvalues[lo_idx.z], cm1_kvalues[lo_idx.w]));
+        pack32(i8vec4(KV_LOOKUP(lo_idx.x), KV_LOOKUP(lo_idx.y),
+                      KV_LOOKUP(lo_idx.z), KV_LOOKUP(lo_idx.w)));
     buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] =
-        pack32(i8vec4(cm1_kvalues[hi_idx.x], cm1_kvalues[hi_idx.y],
-                      cm1_kvalues[hi_idx.z], cm1_kvalues[hi_idx.w]));
+        pack32(i8vec4(KV_LOOKUP(hi_idx.x), KV_LOOKUP(hi_idx.y),
+                      KV_LOOKUP(hi_idx.z), KV_LOOKUP(hi_idx.w)));
 
     if (loadr == 0) {
         buf_a_d[ks * BM + buf_ib] = float(blk.d);
@@ -198,11 +206,11 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
     const u8vec4 lo_idx = unpack8(blk.qs & 0x0F0F0F0F);
     const u8vec4 hi_idx = unpack8((blk.qs >> 4) & 0x0F0F0F0F);
     buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr    ] =
-        pack32(i8vec4(cm1_kvalues[lo_idx.x], cm1_kvalues[lo_idx.y],
-                      cm1_kvalues[lo_idx.z], cm1_kvalues[lo_idx.w]));
+        pack32(i8vec4(KV_LOOKUP(lo_idx.x), KV_LOOKUP(lo_idx.y),
+                      KV_LOOKUP(lo_idx.z), KV_LOOKUP(lo_idx.w)));
     buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] =
-        pack32(i8vec4(cm1_kvalues[hi_idx.x], cm1_kvalues[hi_idx.y],
-                      cm1_kvalues[hi_idx.z], cm1_kvalues[hi_idx.w]));
+        pack32(i8vec4(KV_LOOKUP(hi_idx.x), KV_LOOKUP(hi_idx.y),
+                      KV_LOOKUP(hi_idx.z), KV_LOOKUP(hi_idx.w)));
 
     if (loadr == 0) {
         buf_a_d[ks * BM + buf_ib] = blk.d;
@@ -230,11 +238,11 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
     const u8vec4 lo_idx = unpack8(blk.qs & 0x0F0F0F0F);
     const u8vec4 hi_idx = unpack8((blk.qs >> 4) & 0x0F0F0F0F);
     buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr    ] =
-        pack32(i8vec4(cm1_kvalues[lo_idx.x], cm1_kvalues[lo_idx.y],
-                      cm1_kvalues[lo_idx.z], cm1_kvalues[lo_idx.w]));
+        pack32(i8vec4(KV_LOOKUP(lo_idx.x), KV_LOOKUP(lo_idx.y),
+                      KV_LOOKUP(lo_idx.z), KV_LOOKUP(lo_idx.w)));
     buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + loadr + 4] =
-        pack32(i8vec4(cm1_kvalues[hi_idx.x], cm1_kvalues[hi_idx.y],
-                      cm1_kvalues[hi_idx.z], cm1_kvalues[hi_idx.w]));
+        pack32(i8vec4(KV_LOOKUP(hi_idx.x), KV_LOOKUP(hi_idx.y),
+                      KV_LOOKUP(hi_idx.z), KV_LOOKUP(hi_idx.w)));
 
     if (loadr == 0) {
         buf_a_d[ks * BM + buf_ib] = e8m0_to_fp32(blk.e) * 0.5;
@@ -492,11 +500,11 @@ void block_a_to_shmem(block_a_prefetch blk, uint buf_ib, uint ks, uint loadr) {
     const uint sub_base = (loadr >> 1) * 4;
     const uint byte_group = loadr & 1u;
     buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + sub_base + byte_group] =
-        pack32(i8vec4(cm1_kvalues[lo_idx.x], cm1_kvalues[lo_idx.y],
-                      cm1_kvalues[lo_idx.z], cm1_kvalues[lo_idx.w]));
+        pack32(i8vec4(KV_LOOKUP(lo_idx.x), KV_LOOKUP(lo_idx.y),
+                      KV_LOOKUP(lo_idx.z), KV_LOOKUP(lo_idx.w)));
     buf_a_qs[buf_ib * QPITCH + ks * (BK / 4) + sub_base + 2 + byte_group] =
-        pack32(i8vec4(cm1_kvalues[hi_idx.x], cm1_kvalues[hi_idx.y],
-                      cm1_kvalues[hi_idx.z], cm1_kvalues[hi_idx.w]));
+        pack32(i8vec4(KV_LOOKUP(hi_idx.x), KV_LOOKUP(hi_idx.y),
+                      KV_LOOKUP(hi_idx.z), KV_LOOKUP(hi_idx.w)));
 
     if (loadr == 0) {
         buf_a_d[(ks * KSCALES    ) * BM + buf_ib] = ue4m3_to_fp32(blk.d0) * 0.5;
-- 
2.55.0

