From af52d94b8ae8f5fb476b47da646d193169889b8d Mon Sep 17 00:00:00 2001
From: Mike Lothian <mike@fireburn.co.uk>
Date: Wed, 16 Sep 2026 15:38:10 +0100
Subject: [PATCH 01/18] vulkan : handle IQ4_XS mat-vec with more than 2 columns
 efficiently

Assisted-by: Claude Code (Claude Opus 5.5)
---
 .../vulkan-shaders/mul_mat_vec_iq4_xs.comp    | 37 ++++++++++---------
 1 file changed, 19 insertions(+), 18 deletions(-)

diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iq4_xs.comp b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iq4_xs.comp
index a2b99d9ab..c102177df 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iq4_xs.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/mul_mat_vec_iq4_xs.comp
@@ -9,9 +9,15 @@ layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;
 FLOAT_TYPE temp[NUM_COLS][NUM_ROWS];
 
 // dedicated iq4_xs mat-vec, mirrors mul_mat_vec_iq3_s.comp
-// one packed32 word per l, so the 6-bit subblock scale is hoisted to a single fma after register accumulation
+// one packed32 word per l
 
-void calc_superblock(const uint a_offset, const uint b_offset, const uint ib32, const uint i, const uint num_blocks_per_row, const uint first_row, const uint num_rows) {
+// Invocations per superblock. With more than 2 columns, 8 invocations spill registers, so use 16 or 32.
+const uint TPB = NUM_COLS <= 2 ? 8 : (NUM_COLS <= 4 ? 16 : 32);
+const uint NL  = 32 / TPB; // packed32 words per invocation
+
+void calc_superblock(const uint a_offset, const uint b_offset, const uint itid, const uint i, const uint num_blocks_per_row, const uint first_row, const uint num_rows) {
+    const uint ib32 = itid / (TPB / 8);
+    const uint l0   = (itid % (TPB / 8)) * NL;
     const uint y_idx = i * QUANT_K + 32 * ib32;
 
     uint ibi = a_offset + first_row * num_blocks_per_row + i;
@@ -21,12 +27,8 @@ void calc_superblock(const uint a_offset, const uint b_offset, const uint ib32,
         const uint sh = (data_a[ibi].scales_h >> (2 * ib32)) & 3;
         const float dscale = d * float(int(sl | (sh << 4)) - 32);
 
-        FLOAT_TYPE sum[NUM_COLS];
-        [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) {
-            sum[j] = FLOAT_TYPE(0);
-        }
-
-        [[unroll]] for (uint l = 0; l < 4; ++l) {
+        [[unroll]] for (uint ll = 0; ll < NL; ++ll) {
+            const uint l = l0 + ll;
             const uint w = data_a_packed32[ibi].qs[4 * ib32 + l];
             const u8vec4 q0 = unpack8(w & 0x0F0F0F0F);
             const u8vec4 q1 = unpack8((w >> 4) & 0x0F0F0F0F);
@@ -35,7 +37,8 @@ void calc_superblock(const uint a_offset, const uint b_offset, const uint ib32,
                 const vec4 b0 = vec4(data_b_v4[(j*p.batch_stride_b + b_offset + y_idx) / 4 + l]);
                 const vec4 b1 = vec4(data_b_v4[(j*p.batch_stride_b + b_offset + y_idx) / 4 + 4 + l]);
 
-                sum[j] = fma(FLOAT_TYPE(b0.x), FLOAT_TYPE(kvalues_iq4nl[q0.x]),
+                const FLOAT_TYPE sum =
+                        fma(FLOAT_TYPE(b0.x), FLOAT_TYPE(kvalues_iq4nl[q0.x]),
                         fma(FLOAT_TYPE(b0.y), FLOAT_TYPE(kvalues_iq4nl[q0.y]),
                         fma(FLOAT_TYPE(b0.z), FLOAT_TYPE(kvalues_iq4nl[q0.z]),
                         fma(FLOAT_TYPE(b0.w), FLOAT_TYPE(kvalues_iq4nl[q0.w]),
@@ -43,12 +46,10 @@ void calc_superblock(const uint a_offset, const uint b_offset, const uint ib32,
                         fma(FLOAT_TYPE(b1.y), FLOAT_TYPE(kvalues_iq4nl[q1.y]),
                         fma(FLOAT_TYPE(b1.z), FLOAT_TYPE(kvalues_iq4nl[q1.z]),
                         fma(FLOAT_TYPE(b1.w), FLOAT_TYPE(kvalues_iq4nl[q1.w]),
-                        sum[j]))))))));
-            }
-        }
+                        FLOAT_TYPE(0)))))))));
 
-        [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) {
-            temp[j][n] = fma(dscale, sum[j], temp[j][n]);
+                temp[j][n] = fma(dscale, sum, temp[j][n]);
+            }
         }
 
         ibi += num_blocks_per_row;
@@ -62,11 +63,11 @@ void compute_outputs(const uint32_t first_row, const uint32_t num_rows) {
 
     const uint num_blocks_per_row = p.ncols / QUANT_K;
 
-    // 8 threads are used to process each block
-    const uint blocks_per_wg = gl_WorkGroupSize.x/8;
+    // TPB invocations are used to process each block
+    const uint blocks_per_wg = gl_WorkGroupSize.x/TPB;
     const uint tid = gl_LocalInvocationID.x;
-    const uint itid = tid % 8;  // 0...7
-    const uint ix = tid / 8;
+    const uint itid = tid % TPB;
+    const uint ix = tid / TPB;
 
     [[unroll]] for (uint j = 0; j < NUM_COLS; ++j) {
         [[unroll]] for (uint i = 0; i < NUM_ROWS; ++i) {
-- 
2.55.0

