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));