Commit 1b43d3116 for llama.cpp
commit 1b43d311698269073fb4827db10f9d6f195d2e7d
Author: Pascal <admin@serveurperso.com>
Date: Mon Oct 5 08:39:47 2026 +0200
cuda: tile the lightning indexer over keys and tokens for 4 heads (#29901)
* cuda: tile the lightning indexer over keys and tokens for 4 heads
With too few heads for a wmma tile, a block scores 64 keys against 8
tokens: the keys are staged once in half precision, the queries one
head at a time, and each thread owns one key for two tokens, so no dot
product needs a cross thread reduction. Batches smaller than a token
tile keep the vector kernel. test-backend-ops measures 4 heads.
* cuda: multiply the lightning indexer tile in float
Address review from am17an: the half2 products overflow once a single
q * k exceeds the f16 range. The queries stay in float in shared memory
and each half2 of keys is widened once for both tokens, so every
product and sum is computed in float.
* cuda: widen each lightning indexer key once for all heads
The tile kernel stages the queries and weights of every head at once,
so each key element is widened from half once and feeds all heads,
with a single barrier. F16 keys are copied into the tile without a
float round trip. Keeping the keys in float in shared memory measures
slower, the occupancy drops.
* cuda: stop the lightning indexer tile from spilling registers on ROCm
Each thread of the tile kernel now scores two keys for a single token,
so a warp shares its token and the query reads are broadcasts: six
shared reads per element pair instead of nine for the same products.
The inner loop is unrolled by 8, which keeps gfx908 at 63 VGPRs with no
spill where the fully unrolled loop needed over a thousand, and makes
the kernel 36x faster on an R9700 and slightly faster on CUDA.
diff --git a/ggml/src/ggml-cuda/lightning-indexer.cu b/ggml/src/ggml-cuda/lightning-indexer.cu
index 54e0e6314..90c06a1cf 100644
--- a/ggml/src/ggml-cuda/lightning-indexer.cu
+++ b/ggml/src/ggml-cuda/lightning-indexer.cu
@@ -236,6 +236,9 @@ static __global__ void lightning_indexer_kernel_wmma(
#endif // defined(TURING_MMA_AVAILABLE)
#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
+// tokens scored per block by the tile kernel
+#define LIGHTNING_INDEXER_TILE_TOKENS 8
+
// TODO there is one ugly assumption used in this kernel - that WARP_SIZE is equal to 32
// thanks to that one warp operating on float4 processes whole indexer K/Q vectors
// 32 * 4 = 128 (N_EMBD)
@@ -382,6 +385,152 @@ static __global__ void lightning_indexer_kernel_vec(
}
}
+// one block scores a tile of K_VECS_PER_BLOCK keys against TOKENS_PER_BLOCK tokens: the keys are
+// staged in half precision and the queries of every head in float, each thread owns KEYS_PER_THREAD
+// keys for one token, a warp shares its token so the query reads are broadcasts, and every key
+// element is widened once for all heads, so no dot product needs a cross thread reduction
+template <int WARPS_PER_BLOCK, int K_VECS_PER_BLOCK, int64_t N_EMBD, int64_t N_HEAD, ggml_type TYPE_K>
+static __global__ void lightning_indexer_kernel_tile(
+ const float * Q, const char * K, const float * W, const half * M, float * dst,
+ int64_t n_stream, int64_t n_batch, int64_t n_kv,
+ size_t nb1, size_t nb2, size_t nb3,
+ size_t nbq1, size_t nbq2, size_t nbq3,
+ size_t nbk1, size_t nbk2, size_t nbk3,
+ size_t nbw1, size_t nbw2, size_t nbw3,
+ size_t nbm1, size_t nbm2, size_t nbm3,
+ int64_t nem3
+ ) {
+
+ constexpr int THREADS_PER_BLOCK = WARPS_PER_BLOCK * WARP_SIZE;
+ constexpr int TOKENS_PER_BLOCK = LIGHTNING_INDEXER_TILE_TOKENS;
+ constexpr int KEY_LANES = THREADS_PER_BLOCK / TOKENS_PER_BLOCK;
+ constexpr int KEYS_PER_THREAD = K_VECS_PER_BLOCK / KEY_LANES;
+ constexpr int N_EMBD_H2 = N_EMBD / 2;
+
+ static_assert(THREADS_PER_BLOCK % TOKENS_PER_BLOCK == 0, "threads must cover the token tile");
+ static_assert(K_VECS_PER_BLOCK % KEY_LANES == 0, "key lanes must cover the key tile");
+
+ const int tid = threadIdx.y * WARP_SIZE + threadIdx.x;
+ const int start_kv = blockIdx.x * K_VECS_PER_BLOCK;
+ const int start_batch = blockIdx.y * TOKENS_PER_BLOCK;
+ const int i_stream = blockIdx.z;
+
+ // the row padding keeps the keys of consecutive threads in distinct banks
+ __shared__ half2 k_shared[K_VECS_PER_BLOCK][N_EMBD_H2 + 1];
+ __shared__ float2 q_shared[N_HEAD][TOKENS_PER_BLOCK][N_EMBD_H2];
+ __shared__ float w_shared[N_HEAD][TOKENS_PER_BLOCK];
+
+ // phase 1 - stage the key tile four elements at a time, rows past n_kv are zero
+
+#pragma unroll
+ for (int i = tid; i < K_VECS_PER_BLOCK * (N_EMBD / 4); i += THREADS_PER_BLOCK) {
+ const int r = i / (N_EMBD / 4);
+ const int c4 = i % (N_EMBD / 4);
+
+ half2 lo = make_half2(0.0f, 0.0f);
+ half2 hi = lo;
+ if (start_kv + r < n_kv) {
+ const char * k_row = K + (start_kv + r)*nbk2 + i_stream*nbk3;
+ if constexpr (TYPE_K == GGML_TYPE_F16) {
+ lo = ((const half2 *) k_row)[2*c4 + 0];
+ hi = ((const half2 *) k_row)[2*c4 + 1];
+ } else {
+ float4 v;
+ if constexpr (TYPE_K == GGML_TYPE_F32) {
+ v = ((const float4 *) k_row)[c4];
+ } else {
+ constexpr dequantize_V_t dequantize_k = get_dequantize_V<TYPE_K, float, 4>();
+ dequantize_k(k_row, &v, c4 * 4);
+ }
+ lo = make_half2(v.x, v.y);
+ hi = make_half2(v.z, v.w);
+ }
+ }
+
+ k_shared[r][2*c4 + 0] = lo;
+ k_shared[r][2*c4 + 1] = hi;
+ }
+
+ // phase 2 - stage the queries and weights of every head, tokens past n_batch are zero
+
+#pragma unroll
+ for (int i = tid; i < N_HEAD * TOKENS_PER_BLOCK * (N_EMBD / 4); i += THREADS_PER_BLOCK) {
+ const int h = i / (TOKENS_PER_BLOCK * (N_EMBD / 4));
+ const int r = i / (N_EMBD / 4) % TOKENS_PER_BLOCK;
+ const int c4 = i % (N_EMBD / 4);
+
+ float4 v = make_float4(0.0f, 0.0f, 0.0f, 0.0f);
+ if (start_batch + r < n_batch) {
+ v = *(const float4 *) ((const char *) Q + h*nbq1 + (start_batch + r)*nbq2 + i_stream*nbq3 + c4*sizeof(float4));
+ }
+
+ q_shared[h][r][2*c4 + 0] = make_float2(v.x, v.y);
+ q_shared[h][r][2*c4 + 1] = make_float2(v.z, v.w);
+ }
+
+ if (tid < N_HEAD * TOKENS_PER_BLOCK) {
+ const int h = tid / TOKENS_PER_BLOCK;
+ const int r = tid % TOKENS_PER_BLOCK;
+ w_shared[h][r] = start_batch + r < n_batch ?
+ ((const float *) ((const char *) W + (start_batch + r)*nbw1 + i_stream*nbw3))[h] : 0.0f;
+ }
+
+ __syncthreads();
+
+ // phase 3 - float products of the widened keys for every head, ReLU, weight
+
+ const int kl = tid % KEY_LANES;
+ const int tl = tid / KEY_LANES;
+
+ float qk[N_HEAD][KEYS_PER_THREAD] = { { 0.0f } };
+
+#pragma unroll 8
+ for (int c = 0; c < N_EMBD_H2; ++c) {
+ float2 k_val[KEYS_PER_THREAD];
+#pragma unroll
+ for (int j = 0; j < KEYS_PER_THREAD; ++j) {
+ k_val[j] = __half22float2(k_shared[kl + j*KEY_LANES][c]);
+ }
+#pragma unroll
+ for (int h = 0; h < N_HEAD; ++h) {
+ const float2 q_val = q_shared[h][tl][c];
+#pragma unroll
+ for (int j = 0; j < KEYS_PER_THREAD; ++j) {
+ qk[h][j] = fmaf(k_val[j].x, q_val.x, qk[h][j]);
+ qk[h][j] = fmaf(k_val[j].y, q_val.y, qk[h][j]);
+ }
+ }
+ }
+
+ float score[KEYS_PER_THREAD] = { 0.0f };
+
+#pragma unroll
+ for (int h = 0; h < N_HEAD; ++h) {
+#pragma unroll
+ for (int j = 0; j < KEYS_PER_THREAD; ++j) {
+ score[j] += fmaxf(qk[h][j], 0.0f) * w_shared[h][tl];
+ }
+ }
+
+ // phase 4 - add the mask and write, consecutive threads write consecutive keys
+
+ const int i_batch = start_batch + tl;
+ if (i_batch >= n_batch) {
+ return;
+ }
+
+ const half * m_base = (const half *) ((const char *) M + i_batch*nbm1 + (i_stream%nem3)*nbm3);
+ float * dst_base = (float *) ((char *) dst + i_batch*nb1 + i_stream*nb3);
+
+#pragma unroll
+ for (int j = 0; j < KEYS_PER_THREAD; ++j) {
+ const int i_kv = start_kv + kl + j*KEY_LANES;
+ if (i_kv < n_kv) {
+ dst_base[i_kv] = score[j] + __half2float(m_base[i_kv]);
+ }
+ }
+}
+
#define LIGHTNING_INDEXER_CASE(lightning_indexer_kernel, n_embd, n_head, K, type_K) \
if (K->type == (type_K)) { \
lightning_indexer_kernel<WARPS_PER_BLOCK, K_VECS_PER_BLOCK, n_embd, n_head, type_K> \
@@ -528,8 +677,27 @@ 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 && n_batch >= LIGHTNING_INDEXER_TILE_TOKENS) {
+ // too few heads for a wmma tile, the tile kernel shares the keys across the tokens
+ constexpr int WARPS_PER_BLOCK = 8;
+ constexpr int K_VECS_PER_BLOCK = 64;
+
+ dim3 block(32, WARPS_PER_BLOCK);
+ int num_kv_blocks = (n_kv + (K_VECS_PER_BLOCK) - 1) / (K_VECS_PER_BLOCK);
+ int num_batch_blocks = (n_batch + LIGHTNING_INDEXER_TILE_TOKENS - 1) / LIGHTNING_INDEXER_TILE_TOKENS;
+ dim3 grid(num_kv_blocks, num_batch_blocks, n_stream);
+
+ LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_tile, 128, 4, k, GGML_TYPE_F16)
+ LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_tile, 128, 4, k, GGML_TYPE_Q4_0)
+ LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_tile, 128, 4, k, GGML_TYPE_Q4_1)
+ LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_tile, 128, 4, k, GGML_TYPE_Q5_0)
+ LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_tile, 128, 4, k, GGML_TYPE_Q5_1)
+ LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_tile, 128, 4, k, GGML_TYPE_Q8_0)
+ LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_tile, 128, 4, k, GGML_TYPE_BF16)
+ LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_tile, 128, 4, 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
+ // a batch smaller than a token 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;
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index c82af18ff..bc32aae77 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -12042,7 +12042,7 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_perf() {
// lightning_indexer
for (int kv : { 256, 4096, 65536 }) {
for (int bs : { 1, 512, 2048 }) {
- for (int nh : { 32, 64 }) {
+ for (int nh : { 4, 32, 64 }) {
for (int ns : { 1, 4 }) {
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, ns, type_K));