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;