Commit b809b886d for llama.cpp
commit b809b886d94107349f4e2b1a0d4713d8566565aa
Author: Pascal <admin@serveurperso.com>
Date: Mon Oct 5 15:09:53 2026 +0200
cuda: use the vector lightning indexer kernel on MUSA (#29990)
* cuda: stage the lightning indexer queries in head passes for MUSA
MUSA archs 21 and 22 cap static shared memory at 28 KB, and the tile
kernel staged the queries of all four heads next to the key tile for
33 KB. The queries are now staged in passes of
LIGHTNING_INDEXER_TILE_HEADS_PER_PASS heads: two on MUSA for 25 KB,
four elsewhere where the single pass folds to the previous kernel.
* cuda: use the vector lightning indexer kernel on MUSA
Address review from am17an: the tile kernel stays off MUSA, whose archs
21 and 22 cap static shared memory at 28 KB, below the 33 KB the tile
needs, so MUSA keeps the vector kernel it ran before. This replaces the
head passes, CUDA and ROCm run the merged kernel unchanged.
diff --git a/ggml/src/ggml-cuda/lightning-indexer.cu b/ggml/src/ggml-cuda/lightning-indexer.cu
index 90c06a1cf..57cb1c1dc 100644
--- a/ggml/src/ggml-cuda/lightning-indexer.cu
+++ b/ggml/src/ggml-cuda/lightning-indexer.cu
@@ -677,6 +677,8 @@ 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;
@@ -696,8 +698,9 @@ 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, use vector kernel
+ // a batch smaller than a token tile, or MUSA, 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;