Commit ad2156533 for llama.cpp

commit ad2156533102a0d3c4e5fbdf422dc25fba4d03ba
Author: R0CKSTAR <yeahdongcn@gmail.com>
Date:   Wed Oct 7 17:22:56 2026 +0800

    musa: use the tile lightning indexer kernel (#30080)

    Signed-off-by: Xiaodong Ye <xiaodong.ye@mthreads.com>

diff --git a/ggml/src/ggml-cuda/lightning-indexer.cu b/ggml/src/ggml-cuda/lightning-indexer.cu
index 57cb1c1dc..429d2e200 100644
--- a/ggml/src/ggml-cuda/lightning-indexer.cu
+++ b/ggml/src/ggml-cuda/lightning-indexer.cu
@@ -239,6 +239,14 @@ static __global__ void lightning_indexer_kernel_wmma(
 // tokens scored per block by the tile kernel
 #define LIGHTNING_INDEXER_TILE_TOKENS 8

+// heads whose queries the tile kernel stages per pass, MUSA arch 21 caps static shared memory
+// at 28 KB and the queries of four heads do not fit there next to the key tile
+#if defined(GGML_USE_MUSA) && defined(__MUSA_ARCH__) && __MUSA_ARCH__ < 220
+#define LIGHTNING_INDEXER_TILE_HEADS_PER_PASS 2
+#else
+#define LIGHTNING_INDEXER_TILE_HEADS_PER_PASS 4
+#endif
+
 // 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)
@@ -406,9 +414,11 @@ static __global__ void lightning_indexer_kernel_tile(
     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;
+    constexpr int HEADS_PER_PASS    = N_HEAD < LIGHTNING_INDEXER_TILE_HEADS_PER_PASS ? N_HEAD : LIGHTNING_INDEXER_TILE_HEADS_PER_PASS;

     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");
+    static_assert(N_HEAD % HEADS_PER_PASS == 0, "head passes must cover the heads");

     const int tid         = threadIdx.y * WARP_SIZE + threadIdx.x;
     const int start_kv    = blockIdx.x * K_VECS_PER_BLOCK;
@@ -417,7 +427,7 @@ static __global__ void lightning_indexer_kernel_tile(

     // 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__ float2 q_shared[HEADS_PER_PASS][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
@@ -451,22 +461,7 @@ static __global__ void lightning_indexer_kernel_tile(
         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);
-    }
+    // phase 2 - stage the weights of every head, tokens past n_batch are zero

     if (tid < N_HEAD * TOKENS_PER_BLOCK) {
         const int h = tid / TOKENS_PER_BLOCK;
@@ -475,33 +470,60 @@ static __global__ void lightning_indexer_kernel_tile(
             ((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]);
+    for (int h0 = 0; h0 < N_HEAD; h0 += HEADS_PER_PASS) {
+        // the previous pass is fully consumed before its queries are replaced
+        if (h0 > 0) {
+            __syncthreads();
         }
+
+        // phase 3 - stage the queries of the heads of this pass, tokens past n_batch are zero
+
 #pragma unroll
-        for (int h = 0; h < N_HEAD; ++h) {
-            const float2 q_val = q_shared[h][tl][c];
+        for (int i = tid; i < HEADS_PER_PASS * 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 + (h0 + 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);
+        }
+
+        __syncthreads();
+
+        // phase 4 - float products of the widened keys for the heads of this pass
+
+#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) {
-                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]);
+                k_val[j] = __half22float2(k_shared[kl + j*KEY_LANES][c]);
+            }
+#pragma unroll
+            for (int h = 0; h < HEADS_PER_PASS; ++h) {
+                const float2 q_val = q_shared[h][tl][c];
+#pragma unroll
+                for (int j = 0; j < KEYS_PER_THREAD; ++j) {
+                    qk[h0 + h][j] = fmaf(k_val[j].x, q_val.x, qk[h0 + h][j]);
+                    qk[h0 + h][j] = fmaf(k_val[j].y, q_val.y, qk[h0 + h][j]);
+                }
             }
         }
     }

+    // phase 5 - ReLU, weight, add the mask and write, consecutive threads write consecutive keys
+
     float score[KEYS_PER_THREAD] = { 0.0f };

 #pragma unroll
@@ -512,8 +534,6 @@ static __global__ void lightning_indexer_kernel_tile(
         }
     }

-    // 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;
@@ -677,8 +697,6 @@ 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");
         }
-#ifndef GGML_USE_MUSA
-    // MUSA archs 21 and 22 cap static shared memory at 28 KB, below what the tile kernel stages
     } 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;
@@ -698,9 +716,8 @@ void ggml_cuda_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor *
         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");
-#endif // GGML_USE_MUSA
     } else if (n_embd == 128 && n_head == 4) {
-        // a batch smaller than a token tile, or MUSA, 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;