Commit 889edf43d for llama.cpp

commit 889edf43ddae0cfe9a4564a882764dc879759870
Author: Pascal <admin@serveurperso.com>
Date:   Sat Oct 3 07:19:00 2026 +0200

    qwen4exp : halve the indexer score memory (#29825)

    * qwen4exp : halve the indexer score memory

    The indexer scored all heads in one product and rectified a copy of it,
    so two [n_pool, n_idx_h, n_tokens] f32 tensors were live at once, the
    largest buffers of the graph at long context. Each head now gets its
    own product, rectified and summed in place into one [n_pool, n_tokens]
    score.

    * qwen4exp: let the allocator reuse the indexer score buffers

    Address review from CISC: use plain ggml_add and ggml_relu in the
    indexer head loop. The graph allocator already runs them in place when
    their source has no other consumer, so the _inplace variants are not
    needed. The compute buffer and the speed are unchanged.

    * cuda: support 4 heads in the lightning indexer

    Dispatch 4 heads to the vector kernel, too few for a wmma tile, and
    accept them in supports_op. test-backend-ops covers 4 heads.

    * metal: take the lightning indexer head count as a function constant

    The kernel reads the head count from a function constant and zero fills
    the last head tile, so any head count runs and 64 heads is unchanged.

    * qwen4exp: compute the indexer score with the lightning indexer

    Address review from am17an: the unweighted sum of the rectified head
    scores scaled by 1/sqrt(head_dim) is the lightning indexer with every
    head weight set to that scale, so the indexer calls
    ggml_lightning_indexer on the pooled keys with an f16 pool mask. The
    keys are read once for all heads and no per head score is
    materialized.

    * vulkan: tile the lightning indexer over keys and tokens

    A workgroup scores 64 keys against 8 tokens: the keys are staged once
    in shared memory, the queries one head at a time, and each invocation
    owns one key for two tokens, so no dot product needs a cross invocation
    reduction. The subgroup variant and the flat dispatch are gone, the grid
    is keys x tokens x streams.

    * vectorize vulkan loads and use fp16 dot product

    ---------

    Co-authored-by: Ruben Ortlam <rortlam@redhat.com>

diff --git a/ggml/src/ggml-cuda/lightning-indexer.cu b/ggml/src/ggml-cuda/lightning-indexer.cu
index 5edc967e0..54e0e6314 100644
--- a/ggml/src/ggml-cuda/lightning-indexer.cu
+++ b/ggml/src/ggml-cuda/lightning-indexer.cu
@@ -528,6 +528,25 @@ void ggml_cuda_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor *
             LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, k, GGML_TYPE_F32)
             GGML_ABORT("fatal error");
         }
+    } else if (n_embd == 128 && n_head == 4) {
+        // too few heads for a wmma tile, use vector kernel
+        constexpr int K_VECS_PER_WARP = 8;
+        constexpr int WARPS_PER_BLOCK = 8;
+        constexpr int K_VECS_PER_BLOCK = K_VECS_PER_WARP * WARPS_PER_BLOCK;
+
+        dim3 block(32, WARPS_PER_BLOCK);
+        int num_kv_blocks = (n_kv + (K_VECS_PER_BLOCK) - 1) / (K_VECS_PER_BLOCK);
+        dim3 grid(num_kv_blocks, n_batch, n_stream);
+
+        LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 4, k, GGML_TYPE_F16)
+        LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 4, k, GGML_TYPE_Q4_0)
+        LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 4, k, GGML_TYPE_Q4_1)
+        LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 4, k, GGML_TYPE_Q5_0)
+        LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 4, k, GGML_TYPE_Q5_1)
+        LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 4, k, GGML_TYPE_Q8_0)
+        LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 4, k, GGML_TYPE_BF16)
+        LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 4, k, GGML_TYPE_F32)
+        GGML_ABORT("fatal error");
     } else {
         GGML_ABORT("fatal error");
     }
@@ -556,7 +575,7 @@ bool ggml_cuda_lightning_indexer_supported(int device, const ggml_tensor * dst)
         return false;
     }

-    if (neq1 != 64 && neq1 != 32) {
+    if (neq1 != 64 && neq1 != 32 && neq1 != 4) {
         return false;
     }

diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp
index 8cf2c8212..91b6ef1dc 100644
--- a/ggml/src/ggml-metal/ggml-metal-device.cpp
+++ b/ggml/src/ggml-metal/ggml-metal-device.cpp
@@ -484,13 +484,23 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning_indexe
         const ggml_tensor * op) {
     GGML_ASSERT(op->op == GGML_OP_LIGHTNING_INDEXER);

+    char base[256];
     char name[256];

-    snprintf(name, 256, "kernel_lightning_indexer_%s", ggml_type_name(op->src[1]->type));
+    const int16_t nh = op->src[0]->ne[1];
+
+    snprintf(base, 256, "kernel_lightning_indexer_%s", ggml_type_name(op->src[1]->type));
+    snprintf(name, 256, "%s_nh=%d", base, nh);

     ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name);
     if (!res.pipeline) {
-        res = ggml_metal_library_compile_pipeline(lib, name, name, nullptr);
+        ggml_metal_cv_t cv = ggml_metal_cv_init();
+
+        ggml_metal_cv_set_int16(cv, nh, FC_LIGHTNING_INDEXER + 0);
+
+        res = ggml_metal_library_compile_pipeline(lib, base, name, cv);
+
+        ggml_metal_cv_free(cv);
     }

     return res;
diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m
index 8a74550d1..951cb802a 100644
--- a/ggml/src/ggml-metal/ggml-metal-device.m
+++ b/ggml/src/ggml-metal/ggml-metal-device.m
@@ -1770,8 +1770,7 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te
             }
             return has_simdgroup_mm; // TODO: over-restricted for vec-kernels
         case GGML_OP_LIGHTNING_INDEXER:
-            if (op->src[0]->ne[0] != OP_LIGHTNING_INDEXER_DK ||
-                op->src[0]->ne[1] != OP_LIGHTNING_INDEXER_NH) {
+            if (op->src[0]->ne[0] != OP_LIGHTNING_INDEXER_DK) {
                 return false;
             }
             if (!has_simdgroup_mm ||
diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h
index 3a34c81a4..a5bc79f57 100644
--- a/ggml/src/ggml-metal/ggml-metal-impl.h
+++ b/ggml/src/ggml-metal/ggml-metal-impl.h
@@ -122,6 +122,7 @@
 #define FC_DSV4_HC                     2000
 #define FC_PAD                         2100
 #define FC_FLASH_ATTN_EXT_TENSOR       2200
+#define FC_LIGHTNING_INDEXER           2200

 // op-specific constants
 #define OP_FLASH_ATTN_EXT_NQPSG 8
@@ -136,7 +137,6 @@
 #define OP_FLASH_ATTN_EXT_VEC_NCPSG 32

 #define OP_LIGHTNING_INDEXER_DK    128
-#define OP_LIGHTNING_INDEXER_NH     64
 #define OP_LIGHTNING_INDEXER_NHPTG   8
 #define OP_LIGHTNING_INDEXER_NKPSG   8
 #define OP_LIGHTNING_INDEXER_NSG     8
diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp
index ed4fe47dd..4a7da2d07 100644
--- a/ggml/src/ggml-metal/ggml-metal-ops.cpp
+++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp
@@ -1372,7 +1372,6 @@ int ggml_metal_op_lightning_indexer(ggml_metal_op_t ctx, int idx) {
     GGML_ASSERT(op->type == GGML_TYPE_F32);

     GGML_ASSERT(q->ne[0] == OP_LIGHTNING_INDEXER_DK);
-    GGML_ASSERT(q->ne[1] == OP_LIGHTNING_INDEXER_NH);

     ggml_metal_kargs_lightning_indexer args = {
         /*.n_kv      =*/ (int32_t) k->ne[2],
diff --git a/ggml/src/ggml-metal/kernels/fa_aux.metal b/ggml/src/ggml-metal/kernels/fa_aux.metal
index 89cf0bcd3..ebee1becd 100644
--- a/ggml/src/ggml-metal/kernels/fa_aux.metal
+++ b/ggml/src/ggml-metal/kernels/fa_aux.metal
@@ -329,6 +329,8 @@ kernel void kernel_flash_attn_ext_vec_reduce(
 #undef DV
 }

+constant short FC_lightning_indexer_nh [[function_constant(FC_LIGHTNING_INDEXER + 0)]];
+
 template<
     typename kd4x4_t,
     short nl_k,
@@ -345,7 +347,7 @@ kernel void kernel_lightning_indexer(
         ushort tiisg[[thread_index_in_simdgroup]],
         ushort sgitg[[simdgroup_index_in_threadgroup]]) {
     constexpr short DK    = OP_LIGHTNING_INDEXER_DK;
-    constexpr short NH    = OP_LIGHTNING_INDEXER_NH;
+    const     short NH    = FC_lightning_indexer_nh;
     constexpr short NHPTG = OP_LIGHTNING_INDEXER_NHPTG;
     constexpr short NKPSG = OP_LIGHTNING_INDEXER_NKPSG;
     constexpr short NSG   = OP_LIGHTNING_INDEXER_NSG;
@@ -411,18 +413,22 @@ kernel void kernel_lightning_indexer(
         float score = 0.0f;

         FOR_UNROLL (short i_head = 0; i_head < NH; i_head += NHPTG) {
-            // stage the Q tile [DK, NHPTG] and the (prescaled) head weights
+            // stage the Q tile [DK, NHPTG] and the (prescaled) head weights, heads past NH are zero
             for (short i = tiitg; i < NHPTG*DK4; i += NTG) {
                 const short ih = i/DK4;
                 const short i4 = i%DK4;

-                device const float4 * q4 = (device const float4 *) (pq + (i_head + ih)*args.nbq1);
+                if (i_head + ih < NH) {
+                    device const float4 * q4 = (device const float4 *) (pq + (i_head + ih)*args.nbq1);

-                sq4[ih*DK4 + i4] = half4(q4[i4]);
+                    sq4[ih*DK4 + i4] = half4(q4[i4]);
+                } else {
+                    sq4[ih*DK4 + i4] = half4(0.0h);
+                }
             }

             if (tiitg < NHPTG) {
-                sw[tiitg] = ((device const float *) pw)[i_head + tiitg];
+                sw[tiitg] = i_head + tiitg < NH ? ((device const float *) pw)[i_head + tiitg] : 0.0f;
             }

             threadgroup_barrier(mem_flags::mem_threadgroup);
diff --git a/ggml/src/ggml-vulkan/ggml-vulkan-push-constants.h b/ggml/src/ggml-vulkan/ggml-vulkan-push-constants.h
index 037bcee82..f1f66a628 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan-push-constants.h
+++ b/ggml/src/ggml-vulkan/ggml-vulkan-push-constants.h
@@ -674,9 +674,7 @@ struct vk_op_lightning_indexer_push_constants {
     uint32_t n_kv;
     uint32_t n_heads;
     uint32_t n_tokens;
-    uint32_t n_streams;
     uint32_t n_masks;
-    uint32_t dispatch_x;
     uint32_t q_nb1;
     uint32_t q_nb2;
     uint32_t q_nb3;
diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
index 033c11741..bfc65bed8 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp
+++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
@@ -3654,15 +3654,9 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {

     ggml_vk_create_pipeline(device, device->pipeline_gated_linear_attn_f32, "gated_linear_attn_f32", gated_linear_attn_f32_len, gated_linear_attn_f32_data, "main", 6, sizeof(vk_op_gated_linear_attn_push_constants), {1, 1, 1}, {}, 1);

-    {
-        const bool li_subgroup = device->subgroup_arithmetic && device->subgroup_require_full_support;
-        const size_t li_len   = li_subgroup ? lightning_indexer_subgroup_f32_len  : lightning_indexer_f32_len;
-        const void * li_data  = li_subgroup ? (const void *)lightning_indexer_subgroup_f32_data : (const void *)lightning_indexer_f32_data;
-
-        for (ggml_type k_type : lightning_indexer_k_types) {
-            const std::string name = "lightning_indexer_" + std::string(ggml_type_name(k_type)) + "_k_f32";
-            ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_f32[k_type], name.c_str(), li_len, li_data, "main", 5, sizeof(vk_op_lightning_indexer_push_constants), {1, 1, 1}, {(uint32_t)k_type, fa_block_bytes(k_type), device->subgroup_size}, 1, true, li_subgroup);
-        }
+    for (ggml_type k_type : lightning_indexer_k_types) {
+        const std::string name = "lightning_indexer_" + std::string(ggml_type_name(k_type)) + "_k_f32";
+        ggml_vk_create_pipeline(device, device->pipeline_lightning_indexer_f32[k_type], name.c_str(), lightning_indexer_f32_len, lightning_indexer_f32_data, "main", 5, sizeof(vk_op_lightning_indexer_push_constants), {1, 1, 1}, {(uint32_t)k_type, fa_block_bytes(k_type)}, 1, true);
     }

     {
@@ -10171,9 +10165,9 @@ void ggml_vk_lightning_indexer(ggml_backend_vk_context * ctx, vk_context& subctx
     const uint32_t n_streams = q->ne[3];
     const uint32_t n_masks   = m->ne[3];

-    const uint32_t n_outputs = (uint32_t)(dst->ne[0] * dst->ne[1] * dst->ne[3]);
-    const uint32_t dispatch_x = std::min(n_outputs, ctx->device->properties.limits.maxComputeWorkGroupCount[0]);
-    const uint32_t dispatch_y = CEIL_DIV(n_outputs, dispatch_x);
+    // one workgroup per tile of 64 keys and 8 tokens, see lightning_indexer.comp
+    const uint32_t n_tiles_kv = CEIL_DIV(n_kv, 64);
+    const uint32_t n_tiles_t  = CEIL_DIV(n_tokens, 8);

     // q, w and dst are f32 and m is f16, so their strides are passed in elements;
     // k may be quantized, so its strides stay in bytes
@@ -10190,7 +10184,7 @@ void ggml_vk_lightning_indexer(ggml_backend_vk_context * ctx, vk_context& subctx
     const uint32_t d_nb3 = dst->nb[3] / sizeof(float);

     const vk_op_lightning_indexer_push_constants pc = {
-        n_kv, n_heads, n_tokens, n_streams, n_masks, dispatch_x,
+        n_kv, n_heads, n_tokens, n_masks,
         q_nb1, q_nb2, q_nb3,
         k_nb2, k_nb3,
         w_nb1, w_nb3,
@@ -10200,7 +10194,7 @@ void ggml_vk_lightning_indexer(ggml_backend_vk_context * ctx, vk_context& subctx

     ggml_vk_dispatch_pipeline(ctx, subctx, pipeline,
         {ggml_vk_tensor_subbuffer(ctx, q), ggml_vk_tensor_subbuffer(ctx, k), ggml_vk_tensor_subbuffer(ctx, w), ggml_vk_tensor_subbuffer(ctx, m), ggml_vk_tensor_subbuffer(ctx, dst)},
-        pc, {dispatch_x, dispatch_y, 1});
+        pc, {n_tiles_kv, n_tiles_t, n_streams});
 }

 void ggml_vk_gated_delta_net(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) {
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer.comp b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer.comp
index 9b34d8366..56e8d1862 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/lightning_indexer.comp
@@ -3,10 +3,6 @@
 #extension GL_EXT_control_flow_attributes : require
 #extension GL_EXT_shader_16bit_storage : require
 #extension GL_EXT_shader_explicit_arithmetic_types_float16 : require
-#extension GL_KHR_shader_subgroup_basic : enable
-#if USE_SUBGROUP_ADD
-#extension GL_KHR_shader_subgroup_arithmetic : enable
-#endif

 #define BINDING_IDX_K 0u

@@ -16,14 +12,19 @@

 layout(constant_id = 0) const uint FaTypeK = GGML_TYPE_F32;
 layout(constant_id = 1) const uint FaBlockBytesK = 4;
-layout(constant_id = 2) const uint SUBGROUP_SIZE = 32;

 #include "flash_attn_dequant.glsl"

-// one workgroup computes one output element, one invocation per head element
+// one workgroup scores a tile of BK keys against BT tokens: the keys are staged once,
+// the queries one head at a time, and each invocation owns one key for TPI tokens
 #define HEAD_SIZE 128
+#define WG_SIZE   256
+#define BK        64
+#define BT        8
+#define TS        (WG_SIZE / BK)
+#define TPI       (BT / TS)

-layout(local_size_x = HEAD_SIZE, local_size_y = 1, local_size_z = 1) in;
+layout(local_size_x = WG_SIZE, local_size_y = 1, local_size_z = 1) in;

 layout(binding = 0) readonly buffer QBuf { float q[]; };
 layout(binding = 1) readonly buffer KBufF16 { float16_t k_f16[]; };
@@ -37,9 +38,7 @@ layout(push_constant) uniform PushConstants {
     uint n_kv;
     uint n_heads;
     uint n_tokens;
-    uint n_streams;
     uint n_masks;
-    uint dispatch_x;
     uint q_nb1;
     uint q_nb2;
     uint q_nb3;
@@ -53,99 +52,109 @@ layout(push_constant) uniform PushConstants {
     uint d_nb3;
 };

-shared float k_row[HEAD_SIZE];
-
-#if USE_SUBGROUP_ADD
-shared float sg_partials[HEAD_SIZE / SUBGROUP_SIZE];
-#else
-shared float partials[HEAD_SIZE];
-#endif
+// the row padding keeps the keys of consecutive invocations in distinct banks
+shared f16vec4 k_tile[BK][HEAD_SIZE / 4 + 1];
+shared f16vec4 q_tile[BT][HEAD_SIZE / 4];
+shared float   w_tile[BT];

 void main() {
     const uint tid = gl_LocalInvocationID.x;
-    const uint output_idx = gl_WorkGroupID.y * dispatch_x + gl_WorkGroupID.x;
-    const uint n_outputs = n_kv * n_tokens * n_streams;
+    const uint ik0 = gl_WorkGroupID.x * BK;
+    const uint t0  = gl_WorkGroupID.y * BT;
+    const uint s   = gl_WorkGroupID.z;

     if (fa_type_needs_shmem(FaTypeK)) {
         init_iq_shmem(gl_WorkGroupSize);
     }

-    if (output_idx >= n_outputs) {
-        return;
-    }
-
-    const uint ik = output_idx % n_kv;
-    const uint ts = output_idx / n_kv;
-    const uint t = ts % n_tokens;
-    const uint s = ts / n_tokens;
-    const uint k_offset = ik * k_nb2 + s * k_nb3;
-
     // k strides come in as bytes, so scale them down to the view being indexed
     const uint k_block_elems = fa_block_elems(FaTypeK);
     const uint k_elem_bytes = FaBlockBytesK / k_block_elems;

-    if (FaTypeK == GGML_TYPE_F16) {
-        k_row[tid] = float(k_f16[k_offset / k_elem_bytes + tid]);
-    } else if (FaTypeK == GGML_TYPE_F32) {
-        k_row[tid] = k_f32[k_offset / k_elem_bytes + tid];
-    } else if (FaTypeK == GGML_TYPE_BF16) {
-        k_row[tid] = bf16_to_fp32(uint(k_bf16[k_offset / k_elem_bytes + tid]));
-    } else if (4 * tid < HEAD_SIZE) {
-        const uint coord = 4 * tid;
-        const uint ib = coord / k_block_elems;
-        const uint iqs = coord % k_block_elems;
-        const vec4 values = dequantize4(ib, iqs, k_offset / FaBlockBytesK, BINDING_IDX_K);
-        k_row[coord + 0] = values.x;
-        k_row[coord + 1] = values.y;
-        k_row[coord + 2] = values.z;
-        k_row[coord + 3] = values.w;
+    // stage the key tile four elements at a time, rows past n_kv are zero
+    [[unroll]] for (uint i = tid; i < BK * HEAD_SIZE / 4; i += WG_SIZE) {
+        const uint r = i / (HEAD_SIZE / 4);
+        const uint c4 = i % (HEAD_SIZE / 4);
+
+        vec4 v = vec4(0.0);
+        if (ik0 + r < n_kv) {
+            const uint k_offset = (ik0 + r) * k_nb2 + s * k_nb3;
+            const uint e = k_offset / k_elem_bytes + c4 * 4;
+
+            if (FaTypeK == GGML_TYPE_F16) {
+                v = vec4(k_f16[e], k_f16[e + 1], k_f16[e + 2], k_f16[e + 3]);
+            } else if (FaTypeK == GGML_TYPE_F32) {
+                v = vec4(k_f32[e], k_f32[e + 1], k_f32[e + 2], k_f32[e + 3]);
+            } else if (FaTypeK == GGML_TYPE_BF16) {
+                v = bf16_to_fp32(uvec4(k_bf16[e], k_bf16[e + 1], k_bf16[e + 2], k_bf16[e + 3]));
+            } else {
+                v = dequantize4((c4 * 4) / k_block_elems, (c4 * 4) % k_block_elems, k_offset / FaBlockBytesK, BINDING_IDX_K);
+            }
+        }
+
+        k_tile[r][c4] = f16vec4(v);
     }
-    barrier();

-    const float k_val = k_row[tid];
+    const uint kl = tid % BK;
+    const uint tl = tid / BK;

-    float score = 0.0;
-    for (uint h = 0; h < n_heads; ++h) {
-        const float prod = q[h * q_nb1 + t * q_nb2 + s * q_nb3 + tid] * k_val;
+    float score[TPI];
+    [[unroll]] for (uint j = 0; j < TPI; ++j) {
+        score[j] = 0.0;
+    }

-#if USE_SUBGROUP_ADD
-        const float sg_sum = subgroupAdd(prod);
-        if (gl_SubgroupInvocationID == 0) {
-            sg_partials[gl_SubgroupID] = sg_sum;
-        }
+    for (uint h = 0; h < n_heads; ++h) {
+        // the previous head is fully consumed and, on the first pass, the key tile is complete
         barrier();

-        if (tid == 0) {
-            float sum = 0.0;
-            [[unroll]] for (uint i = 0; i < HEAD_SIZE / SUBGROUP_SIZE; ++i) {
-                sum += sg_partials[i];
+        [[unroll]] for (uint i = tid; i < BT * HEAD_SIZE / 4; i += WG_SIZE) {
+            const uint r = i / (HEAD_SIZE / 4);
+            const uint c4 = i % (HEAD_SIZE / 4);
+            const uint t = t0 + r;
+
+            vec4 v = vec4(0.0);
+            if (t < n_tokens) {
+                const uint q_base = h * q_nb1 + t * q_nb2 + s * q_nb3 + c4 * 4;
+                v = vec4(q[q_base], q[q_base + 1], q[q_base + 2], q[q_base + 3]);
             }
-            score += max(sum, 0.0) * weights[h + t * w_nb1 + s * w_nb3];
+            q_tile[r][c4] = f16vec4(v);
         }
-        // the reads above must complete before the next iteration overwrites sg_partials
-        barrier();
-#else
-        partials[tid] = prod;
+
+        if (tid < BT) {
+            const uint t = t0 + tid;
+            w_tile[tid] = t < n_tokens ? weights[h + t * w_nb1 + s * w_nb3] : 0.0;
+        }
+
         barrier();

-        [[unroll]] for (uint stride = HEAD_SIZE / 2; stride > 0; stride >>= 1) {
-            if (tid < stride) {
-                partials[tid] += partials[tid + stride];
+        float qk[TPI];
+        [[unroll]] for (uint j = 0; j < TPI; ++j) {
+            qk[j] = 0.0;
+        }
+
+        [[unroll]] for (uint c4 = 0; c4 < HEAD_SIZE / 4; ++c4) {
+            const f16vec4 kv = k_tile[kl][c4];
+            [[unroll]] for (uint j = 0; j < TPI; ++j) {
+                const f16vec4 qv = q_tile[tl + j * TS][c4];
+                qk[j] += float(dot(kv, qv));
             }
-            barrier();
         }

-        if (tid == 0) {
-            score += max(partials[0], 0.0) * weights[h + t * w_nb1 + s * w_nb3];
+        [[unroll]] for (uint j = 0; j < TPI; ++j) {
+            score[j] += max(qk[j], 0.0) * w_tile[tl + j * TS];
         }
-        // the read of partials[0] above must complete before the next iteration
-        // overwrites partials[tid]
-        barrier();
-#endif
     }

-    if (tid == 0) {
-        const uint mask_offset = ik + t * m_nb1 + (s % n_masks) * m_nb3;
-        dst[ik + t * d_nb1 + s * d_nb3] = score + float(mask[mask_offset]);
+    const uint ik = ik0 + kl;
+    if (ik >= n_kv) {
+        return;
+    }
+
+    [[unroll]] for (uint j = 0; j < TPI; ++j) {
+        const uint t = t0 + tl + j * TS;
+        if (t < n_tokens) {
+            const uint mask_offset = ik + t * m_nb1 + (s % n_masks) * m_nb3;
+            dst[ik + t * d_nb1 + s * d_nb3] = score[j] + float(mask[mask_offset]);
+        }
     }
 }
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp
index e7e303e50..7597b3d33 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp
@@ -1149,7 +1149,6 @@ void process_shaders() {
     // K quant type is selected at runtime via the FaTypeK spec constant.
     std::map<std::string, std::string> li_dict = {{"FLOAT_TYPE", "float"}, {"FLOAT_TYPEV4", "vec4"}, {"DATA_A_IQ4_NL", "1"}};
     string_to_spv("lightning_indexer_f32", "lightning_indexer.comp", li_dict);
-    string_to_spv("lightning_indexer_subgroup_f32", "lightning_indexer.comp", merge_maps(li_dict, {{"USE_SUBGROUP_ADD", "1"}}));

     string_to_spv("rwkv_wkv7_f32", "wkv7.comp", merge_maps(base_dict, {{"A_TYPE", "float"}}));

diff --git a/src/models/qwen4exp.cpp b/src/models/qwen4exp.cpp
index facab0ebf..f4df6a5c2 100644
--- a/src/models/qwen4exp.cpp
+++ b/src/models/qwen4exp.cpp
@@ -682,7 +682,7 @@ public:
     ggml_tensor * k_idxs        = nullptr; // I64 [n_tokens]
     ggml_tensor * pool_cells    = nullptr; // I32 [n_pool]         cell caching each block's pooled key
     ggml_tensor * pool_idxs     = nullptr; // I32 [kpool, n_pool]  member cells per block, n_kv sentinel for the padded blocks
-    ggml_tensor * pool_mask     = nullptr; // F32 [n_pool, n_tokens]
+    ggml_tensor * pool_mask     = nullptr; // F16 [n_pool, n_tokens]
     ggml_tensor * tail_idxs     = nullptr; // I32 [kpool - 1, n_tokens]
     ggml_tensor * new_pool_idxs = nullptr; // I32 [kpool, n_new]   members of the blocks to re-pool this ubatch
     ggml_tensor * new_pool_rep  = nullptr; // I64 [n_new]          cell to write each new pooled key into
@@ -708,7 +708,7 @@ llama_model_qwen4exp::llm_graph_input_kpool * llama_model_qwen4exp::graph::build
     inp->k_idxs     = mctx_idx->build_input_k_idxs(ctx0, ubatch);
     inp->pool_cells = ggml_new_tensor_1d(ctx0, GGML_TYPE_I32, n_pool);
     inp->pool_idxs  = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, kpool, n_pool);
-    inp->pool_mask  = ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_pool, n_tokens);
+    inp->pool_mask  = ggml_new_tensor_2d(ctx0, GGML_TYPE_F16, n_pool, n_tokens);
     inp->tail_idxs  = ggml_new_tensor_2d(ctx0, GGML_TYPE_I32, kpool - 1, n_tokens);
     ggml_set_input(inp->pool_cells);
     ggml_set_input(inp->pool_idxs);
@@ -812,20 +812,11 @@ ggml_tensor * llama_model_qwen4exp::graph::build_qsa_sel(
             ext_factor, attn_factor, beta_fast, beta_slow);
     cb(q, "indexer_q", il);

-    // the reference sums the rectified head scores unweighted, scaled by 1/sqrt(head_dim)
-    // one product for all heads, then the heads are summed as slices, so nothing is transposed
-    ggml_tensor * kq = ggml_mul_mat(ctx0,
-            ggml_reshape_2d(ctx0, pooled, idx_dim, n_pool),
-            ggml_reshape_2d(ctx0, q, idx_dim, n_idx_h*n_tokens)); // [n_pool, n_idx_h*n_tokens]
-    kq = ggml_relu(ctx0, ggml_reshape_3d(ctx0, kq, n_pool, n_idx_h, n_tokens));
-
-    ggml_tensor * score = nullptr;
-    for (int64_t h = 0; h < n_idx_h; ++h) {
-        ggml_tensor * slice = ggml_view_2d(ctx0, kq, n_pool, n_tokens, kq->nb[2], h*kq->nb[1]);
-        score = score ? ggml_add(ctx0, score, slice) : ggml_cont(ctx0, slice);
-    }
-    score = ggml_scale(ctx0, score, 1.0f/sqrtf((float) idx_dim));
-    score = ggml_add(ctx0, score, inp_kpool->pool_mask); // [n_pool, n_tokens]
+    // the reference sums the rectified head scores unweighted, scaled by 1/sqrt(head_dim),
+    // which is the lightning indexer with every head weight set to that scale
+    ggml_tensor * weights = ggml_fill(ctx0, ggml_new_tensor_2d(ctx0, GGML_TYPE_F32, n_idx_h, n_tokens), 1.0f/sqrtf((float) idx_dim));
+    ggml_tensor * score = ggml_lightning_indexer(ctx0, q, pooled, weights, inp_kpool->pool_mask); // [n_pool, n_tokens]
+    res->add_fused_node({LLM_FUSED_OP_LIGHTNING_INDEXER, score, il});
     cb(score, "indexer_score", il);

     const int64_t n_top_pool = std::min<int64_t>(n_pool, hparams.indexer_top_k / kpool);
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index 98ce061c5..5c7f61a78 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -11417,7 +11417,7 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
     // lightning_indexer
     for (int kv : { 256 }) {
         for (int bs : { 1, 512 }) {
-            for (int nh : { 32, 64 }) {
+            for (int nh : { 4, 32, 64 }) {
                 for (auto [ns, nm] : { std::pair{1, 1}, std::pair{4, 4}, std::pair{4, 1} }) {
                     for (ggml_type type_K : {GGML_TYPE_F32, GGML_TYPE_F16, GGML_TYPE_BF16, GGML_TYPE_Q8_0, GGML_TYPE_Q5_1, GGML_TYPE_Q5_0, GGML_TYPE_Q4_1, GGML_TYPE_Q4_0, GGML_TYPE_IQ4_NL}) {
                         test_cases.emplace_back(new test_lightning_indexer(128, nh, kv, bs, ns, nm, type_K));