]> git.djapps.eu Git - pkg/ggml/sources/llama.cpp/commitdiff
cuda : CUDA GGML_OP_LIGHTNING_INDEXER implementation (generic vector kernel + wmma...
authorfairydreaming <redacted>
Wed, 15 Jul 2026 17:57:52 +0000 (19:57 +0200)
committerGitHub <redacted>
Wed, 15 Jul 2026 17:57:52 +0000 (19:57 +0200)
* cuda : CUDA GGML_OP_LIGHTNING_INDEXER implementation (generic vector kernel + wmma kernel)

* chore : remove indentation of #pragma unroll

* cuda : remove unnecessary kernel template declarations

* cuda : add WARPS_PER_BLOCK and K_VECS_PER_BLOCK template parameters in lightning indexer kernels to avoid duplication of constants.

* cuda : relax MMA architecture requirements to Turing in lightning indexer implementation

* chore : renamed variables

* chore : rename ggml_cuda_op_lightning_indexer() to ggml_cuda_lightning_indexer()

* chore : TODO for AMD rocWMMA

* chore : whitespace formatting

* chore : another variable rename to fix problems caused by shadowing

* chore : yet another rename, this time uppercased all constants

* cuda : added alignment checks for Q and K tensors in lightning indexer implementation

---------

Co-authored-by: Stanisław Szymczyk <redacted>
ggml/src/ggml-cuda/ggml-cuda.cu
ggml/src/ggml-cuda/lightning-indexer.cu [new file with mode: 0644]
ggml/src/ggml-cuda/lightning-indexer.cuh [new file with mode: 0644]

index 0e185e8496eec6cb6d4ea3a551708226efc74be5..e27ab1526fc51732dff780db2bf361276dac65a8 100644 (file)
@@ -65,6 +65,7 @@
 #include "ggml-cuda/tri.cuh"
 #include "ggml-cuda/cumsum.cuh"
 #include "ggml-cuda/fill.cuh"
+#include "ggml-cuda/lightning-indexer.cuh"
 #include "ggml.h"
 
 #include <algorithm>
@@ -2257,6 +2258,9 @@ static bool ggml_cuda_compute_forward(ggml_backend_cuda_context & ctx, struct gg
         case GGML_OP_FILL:
             ggml_cuda_op_fill(ctx, dst);
             break;
+        case GGML_OP_LIGHTNING_INDEXER:
+            ggml_cuda_lightning_indexer(ctx, dst);
+            break;
         default:
             return false;
     }
@@ -4987,6 +4991,8 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
         case GGML_OP_DIAG:
         case GGML_OP_SOLVE_TRI:
             return true;
+        case GGML_OP_LIGHTNING_INDEXER:
+            return ggml_cuda_lightning_indexer_supported(dev_ctx->device, op);
 
         default:
             return false;
diff --git a/ggml/src/ggml-cuda/lightning-indexer.cu b/ggml/src/ggml-cuda/lightning-indexer.cu
new file mode 100644 (file)
index 0000000..5edc967
--- /dev/null
@@ -0,0 +1,588 @@
+#include "common.cuh"
+#include "lightning-indexer.cuh"
+#include "fattn-common.cuh"
+#include "convert.cuh"
+
+#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
+#if defined(TURING_MMA_AVAILABLE)
+
+typedef union {
+    int2 i2;
+    half2 h2[2];
+} half4;
+
+// TODO add support for AMD cards via rocWMMA
+#include <mma.h>
+namespace wmma = nvcuda::wmma;
+
+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_wmma(
+        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 HEADS_PER_INNER_LOOP = 8;
+    constexpr int K_EMBD_PER_INNER_LOOP = 16;
+    constexpr int N_EMBD_PADDED = N_EMBD + 8;
+
+    const int i_batch  = blockIdx.y;
+    const int i_stream = blockIdx.z;
+    const int i_warp   = threadIdx.y;
+    const int i_lane   = threadIdx.x;
+    const int tid      = i_warp * WARP_SIZE + i_lane;
+
+    // each block processes K_VECS_PER_BLOCK K vectors
+    const int start_kv = blockIdx.x * K_VECS_PER_BLOCK;
+
+    const char  * q_base = (const char  *)                 Q + i_batch*nbq2 + i_stream*nbq3;
+    const float * w_base = (const float *) ((const char *) W + i_batch*nbw1 + i_stream*nbw3);
+
+    // phase 1 - load weights and first Q tile to shared memory
+
+    __shared__ float w_shared[N_HEAD];
+    __shared__ int2  q_shared_h[HEADS_PER_INNER_LOOP][N_EMBD_PADDED / 4];
+
+    if (tid < N_HEAD) {
+        w_shared[tid] = w_base[tid];
+    }
+
+    // total number of half4 elements in HEADS_PER_INNER_LOOP x N_EMBD Q tile
+    constexpr int N_Q_TILE = HEADS_PER_INNER_LOOP * (N_EMBD / 4);
+    // number of registers needed in each thread to store Q tile in thread block
+    constexpr int N_Q_NEXT = (N_Q_TILE + THREADS_PER_BLOCK - 1) / THREADS_PER_BLOCK;
+
+#pragma unroll
+    for (int i_q = tid; i_q < N_Q_TILE; i_q += THREADS_PER_BLOCK) {
+        const int i_head = i_q / (N_EMBD / 4);
+        const int i_embd = i_q % (N_EMBD / 4);
+        const float4 q = *(const float4 *) (q_base + i_head*nbq1 + i_embd*sizeof(float4));
+        half4 q_packed;
+        q_packed.h2[0] = __float22half2_rn(make_float2(q.x, q.y));
+        q_packed.h2[1] = __float22half2_rn(make_float2(q.z, q.w));
+        q_shared_h[i_head][i_embd] = q_packed.i2;
+    }
+
+    // phase 2 - load (and dequantize if needed) K to shared mem
+
+    __shared__ half2 k_shared_h[K_VECS_PER_BLOCK][N_EMBD_PADDED / 4][2];
+
+    constexpr int n_k = K_VECS_PER_BLOCK * (N_EMBD / 4);
+
+    if constexpr (TYPE_K == GGML_TYPE_F16) {
+#pragma unroll
+        for (int i_k = tid; i_k < n_k; i_k += THREADS_PER_BLOCK) {
+            const int i_k_vec = i_k / (N_EMBD / 4);
+            const int i_embd = i_k % (N_EMBD / 4);
+            const int i_kv = start_kv + i_k_vec;
+            if (i_kv < n_kv) {
+                const int2 * k_base = (const int2 *) ((const char *) K + i_kv*nbk2 + i_stream*nbk3);
+                *(int2*) &k_shared_h[i_k_vec][i_embd] = k_base[i_embd];
+            } else {
+                *(int2*) &k_shared_h[i_k_vec][i_embd] = make_int2(0, 0);
+            }
+        }
+    } else {
+        constexpr dequantize_V_t dequantize_k = get_dequantize_V<TYPE_K, half, 4>();
+#pragma unroll
+        for (int i_k = tid; i_k < n_k; i_k += THREADS_PER_BLOCK) {
+            const int i_k_vec = i_k / (N_EMBD / 4);
+            const int i_embd = i_k % (N_EMBD / 4);
+            const int i_kv = start_kv + i_k_vec;
+            if (i_kv < n_kv) {
+                const void * k_base = (const void *) ((const char *) K + i_kv*nbk2 + i_stream*nbk3);
+                dequantize_k(k_base, &k_shared_h[i_k_vec][i_embd][0], i_embd * 4);
+            } else {
+                *(int2*) &k_shared_h[i_k_vec][i_embd] = make_int2(0, 0);
+            }
+        }
+    }
+
+    __syncthreads();
+
+    // phase 3 - calculate lightning indexer scores
+
+    __shared__ float qk_shared[WARPS_PER_BLOCK][HEADS_PER_INNER_LOOP][K_VECS_PER_BLOCK];
+
+    // load K fragment
+    wmma::fragment<wmma::matrix_b, HEADS_PER_INNER_LOOP, K_VECS_PER_BLOCK, K_EMBD_PER_INNER_LOOP, half, wmma::col_major> frag_k;
+    wmma::load_matrix_sync(frag_k, (half*) &k_shared_h[0][i_warp * K_EMBD_PER_INNER_LOOP / 4], N_EMBD_PADDED);
+
+    float score_k = 0.0f;
+
+    for (int i_head_0 = 0; i_head_0 < N_HEAD; i_head_0 += HEADS_PER_INNER_LOOP) {
+        const int i_head_next = i_head_0 + HEADS_PER_INNER_LOOP;
+
+        // we don't use accumulator for anything, fill it with zeros
+        wmma::fragment<wmma::accumulator, HEADS_PER_INNER_LOOP, K_VECS_PER_BLOCK, K_EMBD_PER_INNER_LOOP, float> frag_acc;
+        wmma::fill_fragment(frag_acc, 0.0f);
+
+        // load Q fragment
+        wmma::fragment<wmma::matrix_a, HEADS_PER_INNER_LOOP, K_VECS_PER_BLOCK, K_EMBD_PER_INNER_LOOP, half, wmma::row_major> frag_q;
+        wmma::load_matrix_sync(frag_q, (half*) &q_shared_h[0][i_warp * K_EMBD_PER_INNER_LOOP / 4], N_EMBD_PADDED);
+
+        // preload next Q tile to registers during matrix multiplication
+        float4 q_next[N_Q_NEXT];
+
+        if (i_head_next < N_HEAD) {
+#pragma unroll
+            for (int i_q = tid, i_q_next = 0; i_q < N_Q_TILE; i_q += THREADS_PER_BLOCK) {
+                const int i_head = i_head_next + i_q / (N_EMBD / 4);
+                const int i_embd =               i_q % (N_EMBD / 4);
+                q_next[i_q_next++] = *(const float4 *) (q_base + i_head*nbq1 + i_embd*sizeof(float4));
+            }
+        }
+
+        // perform matrix multiplication
+        wmma::mma_sync(frag_acc, frag_q, frag_k, frag_acc);
+        wmma::store_matrix_sync((float*) &qk_shared[i_warp][0][0], frag_acc, K_VECS_PER_BLOCK, wmma::mem_row_major);
+
+        // make sure all threads finished using q_shared_h so we can store next tile
+        __syncthreads();
+
+        // write preloaded Q tile to shared memory
+        if (i_head_next < N_HEAD) {
+#pragma unroll
+            for (int i_q = tid, i_q_next = 0; i_q < N_Q_TILE; i_q += THREADS_PER_BLOCK) {
+                const int i_head = i_q / (N_EMBD / 4);
+                const int i_embd = i_q % (N_EMBD / 4);
+                half4 q_packed;
+                q_packed.h2[0] = __float22half2_rn(make_float2(q_next[i_q_next].x, q_next[i_q_next].y));
+                q_packed.h2[1] = __float22half2_rn(make_float2(q_next[i_q_next].z, q_next[i_q_next].w));
+                q_shared_h[i_head][i_embd] = q_packed.i2;
+                ++i_q_next;
+            }
+        }
+
+        // accumulate QK multiplication results from all block warps
+        // (there are 256 threads in block and 256 matmul outputs)
+        // TODO it will break if WARP_SIZE is not 32
+        const int h = tid / K_VECS_PER_BLOCK;
+        const int k = tid % K_VECS_PER_BLOCK;
+        const float w_val = w_shared[i_head_0 + h];
+
+        float sum = 0.0f;
+#pragma unroll
+        for (int w = 0; w < WARPS_PER_BLOCK; ++w) {
+            sum += qk_shared[w][h][k];
+        }
+
+        // ReLU, weight
+        sum = sum > 0.0f ? sum : 0.0f;
+        sum *= w_val;
+
+        // wait until qk_shared[0] is no longer used
+        __syncthreads();
+
+        // reuse qk_shared[0] for storing partial results
+        qk_shared[0][h][k] = sum;
+
+        // wait until all threads write their results
+        __syncthreads();
+
+        // accumulate result over heads
+        if (tid < K_VECS_PER_BLOCK) {
+#pragma unroll
+            for (int i_head = 0; i_head < HEADS_PER_INNER_LOOP; ++i_head) {
+                score_k += qk_shared[0][i_head][tid];
+            }
+        }
+
+        // make sure all threads finished using qk_shared
+        __syncthreads();
+    }
+
+    // phase 4 - store output to VRAM
+
+    if (tid < K_VECS_PER_BLOCK) {
+        const int i_kv = start_kv + tid;
+        if (i_kv < n_kv) {
+            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);
+            dst_base[i_kv] = score_k + __half2float(m_base[i_kv]);
+        }
+    }
+}
+
+#else // defined(TURING_MMA_AVAILABLE)
+
+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_wmma(
+        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
+    ) {
+    GGML_UNUSED_VARS(Q, K, W, M, dst,
+        n_stream, n_batch, n_kv,
+        nb1, nb2, nb3,
+        nbq1, nbq2, nbq3,
+        nbk1, nbk2, nbk3,
+        nbw1, nbw2, nbw3,
+        nem3);
+    NO_DEVICE_CODE;
+}
+
+#endif // defined(TURING_MMA_AVAILABLE)
+#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
+
+// 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)
+
+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_vec(
+        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 K_VECS_PER_WARP = K_VECS_PER_BLOCK / WARPS_PER_BLOCK;
+    constexpr int THREADS_PER_BLOCK = WARPS_PER_BLOCK * WARP_SIZE;
+
+    const int i_batch  = blockIdx.y;
+    const int i_stream = blockIdx.z;
+    const int i_warp   = threadIdx.y;
+    const int i_lane   = threadIdx.x;
+    const int tid      = i_warp * WARP_SIZE + i_lane;
+
+    // each warp processes K_VECS_PER_WARP K vectors
+    const int start_kv_block = blockIdx.x * K_VECS_PER_BLOCK;
+    const int start_kv = start_kv_block + i_warp * K_VECS_PER_WARP;
+
+    const char  * q_base = (const char  *)                 Q + i_batch*nbq2 + i_stream*nbq3;
+    const float * w_base = (const float *) ((const char *) W + i_batch*nbw1 + i_stream*nbw3);
+
+    // phase 1 - load (and dequantize if needed) K to registers
+
+    float4 k_reg_f[K_VECS_PER_WARP];
+
+    if constexpr (TYPE_K == GGML_TYPE_F32) {
+        // direct copy of float4
+#pragma unroll
+        for (int k = 0; k < K_VECS_PER_WARP; ++k) {
+            int i_kv = start_kv + k;
+            if (i_kv < n_kv) {
+                const float4 * k_base = (const float4 *) ((const char *) K + i_kv*nbk2 + i_stream*nbk3);
+                k_reg_f[k] = k_base[i_lane];
+            } else {
+                k_reg_f[k] = make_float4(0, 0, 0, 0);
+            }
+        }
+    } else {
+        // dequantize remaining types to float
+        constexpr dequantize_V_t dequantize_k = get_dequantize_V<TYPE_K, float, 4>();
+#pragma unroll
+        for (int k = 0; k < K_VECS_PER_WARP; ++k) {
+            int i_kv = start_kv + k;
+            if (i_kv < n_kv) {
+                const void * k_base = (const void *) ((const char *) K + i_kv*nbk2 + i_stream*nbk3);
+                dequantize_k(k_base, &k_reg_f[k], i_lane * 4);
+            } else {
+                k_reg_f[k] = make_float4(0, 0, 0, 0);
+            }
+        }
+    }
+
+    float score_k[K_VECS_PER_WARP] = { 0.0f };
+
+    // load weights and Q only for N_HEAD_INNER heads at once to reduce shared memory usage
+    constexpr int N_HEAD_INNER = N_HEAD / 4;
+
+    for (int i_head_0 = 0; i_head_0 < N_HEAD; i_head_0 += N_HEAD_INNER) {
+        // phase 2 - load weights and Q to shared memory
+
+        __shared__ float  w_shared[N_HEAD_INNER];
+        __shared__ float4 q_shared_f[N_HEAD_INNER][N_EMBD / 4];
+
+        if (tid < N_HEAD_INNER) {
+            w_shared[tid] = w_base[i_head_0 + tid];
+        }
+
+        constexpr int n_q = N_HEAD_INNER * (N_EMBD / 4);
+#pragma unroll
+        for (int i_q = tid; i_q < n_q; i_q += THREADS_PER_BLOCK) {
+            const int i_head_inner = i_q / (N_EMBD / 4);
+            const int i_head = i_head_0 + i_head_inner;
+            const int i_embd = i_q % (N_EMBD / 4);
+            q_shared_f[i_head_inner][i_embd] = *(const float4 *) (q_base + i_head*nbq1 + i_embd*sizeof(float4));
+        }
+
+        __syncthreads();
+
+        // phase 3 - calculate lightning indexer scores
+
+        for (int i_head_inner = 0; i_head_inner < N_HEAD_INNER; ++i_head_inner) {
+            const float w_val = w_shared[i_head_inner];
+            float qk[K_VECS_PER_WARP] = { 0.0f };
+
+            // dot product of floats
+            const float4 q_vec = q_shared_f[i_head_inner][i_lane];
+
+#pragma unroll
+            for (int k = 0; k < K_VECS_PER_WARP; ++k) {
+                ggml_cuda_mad(qk[k], q_vec.x, k_reg_f[k].x);
+                ggml_cuda_mad(qk[k], q_vec.y, k_reg_f[k].y);
+                ggml_cuda_mad(qk[k], q_vec.z, k_reg_f[k].z);
+                ggml_cuda_mad(qk[k], q_vec.w, k_reg_f[k].w);
+            }
+
+#pragma unroll
+            for (int k = 0; k < K_VECS_PER_WARP; ++k) {
+                float sum = warp_reduce_sum(qk[k]);
+
+                // ReLU, weight
+                if (i_lane == 0) {
+                    sum = (sum > 0.0f) ? sum : 0.0f;
+                    score_k[k] += sum * w_val;
+                }
+            }
+        }
+
+        __syncthreads();
+    }
+
+    // phase 4 - store outputs to shared memory
+
+    __shared__ float dst_shared[K_VECS_PER_BLOCK];
+
+    if (i_lane == 0) {
+#pragma unroll
+        for (int k = 0; k < K_VECS_PER_WARP; ++k) {
+            dst_shared[i_warp * K_VECS_PER_WARP + k] = score_k[k];
+        }
+    }
+
+    __syncthreads();
+
+    // phase 5 - write from shared memory to VRAM in coalesced manner
+
+    if (tid < K_VECS_PER_BLOCK) {
+        int i_kv = start_kv_block + tid;
+        if (i_kv < n_kv) {
+            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);
+            dst_base[i_kv] = dst_shared[tid] + __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> \
+            <<<grid, block, 0, ctx.stream()>>>(                                             \
+            q_d, k_d, w_d, m_d, dst_d,                                                      \
+            n_stream, n_batch, n_kv,                                                        \
+            nb1, nb2, nb3,                                                                  \
+            nbq1, nbq2, nbq3,                                                               \
+            nbk1, nbk2, nbk3,                                                               \
+            nbw1, nbw2, nbw3,                                                               \
+            nbm1, nbm2, nbm3,                                                               \
+            nem3                                                                            \
+        );                                                                                  \
+    } else
+
+void ggml_cuda_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
+    const ggml_tensor * q = dst->src[0];
+    const ggml_tensor * k = dst->src[1];
+    const ggml_tensor * w = dst->src[2]; // weights
+    const ggml_tensor * m = dst->src[3]; // mask
+
+    GGML_ASSERT(dst->type == GGML_TYPE_F32);
+    GGML_ASSERT(  q->type == GGML_TYPE_F32);
+    GGML_ASSERT(  w->type == GGML_TYPE_F32);
+    GGML_ASSERT(  m->type == GGML_TYPE_F16);
+
+    GGML_TENSOR_LOCALS(int64_t, neq,  q, ne)
+    GGML_TENSOR_LOCALS(size_t,  nbq,  q, nb)
+    GGML_TENSOR_LOCALS(int64_t, nek,  k, ne)
+    GGML_TENSOR_LOCALS(size_t,  nbk,  k, nb)
+    GGML_TENSOR_LOCALS(int64_t, new,  w, ne)
+    GGML_TENSOR_LOCALS(size_t,  nbw,  w, nb)
+    GGML_TENSOR_LOCALS(int64_t, nem,  m, ne)
+    GGML_TENSOR_LOCALS(size_t,  nbm,  m, nb)
+    GGML_TENSOR_LOCALS(int64_t, ne, dst, ne)
+    GGML_TENSOR_LOCALS(size_t,  nb, dst, nb)
+
+    // input tensor rows must be contiguous
+    GGML_ASSERT(nbq0 == ggml_type_size(q->type));
+    GGML_ASSERT(nbk0 == ggml_type_size(k->type));
+    GGML_ASSERT(nbw0 == ggml_type_size(w->type));
+    GGML_ASSERT(nbm0 == ggml_type_size(m->type));
+
+    // dst cannot be transposed or permuted
+    GGML_ASSERT(nb0 == sizeof(float));
+    GGML_ASSERT(nb0 <= nb1);
+    GGML_ASSERT(nb1 <= nb2);
+    GGML_ASSERT(nb2 <= nb3);
+
+    const int n_embd   = q->ne[0];
+    const int n_head   = q->ne[1];
+    const int n_batch  = q->ne[2];
+    const int n_stream = q->ne[3];
+    const int n_kv     = k->ne[2];
+
+    const float *   q_d = (const float *)   q->data;
+    const char  *   k_d = (const char  *)   k->data;
+    const float *   w_d = (const float *)   w->data;
+    const half  *   m_d = (const half  *)   m->data;
+    float       * dst_d = (      float *) dst->data;
+
+    const int device = ggml_cuda_get_device();
+    const int cc     = ggml_cuda_info().devices[device].cc;
+
+    if (n_embd == 128 && n_head == 64) {
+#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
+        if (GGML_CUDA_CC_IS_NVIDIA(cc) && turing_mma_available(cc) && k->type != GGML_TYPE_F32 && k->type != GGML_TYPE_BF16) {
+            // use wmma kernel
+            constexpr int K_VECS_PER_BLOCK = 32;
+            constexpr int WARPS_PER_BLOCK = 8;
+
+            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_wmma, 128, 64, k, GGML_TYPE_F16)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, k, GGML_TYPE_Q4_0)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, k, GGML_TYPE_Q4_1)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, k, GGML_TYPE_Q5_0)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, k, GGML_TYPE_Q5_1)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, k, GGML_TYPE_Q8_0)
+            GGML_ABORT("fatal error");
+        } else {
+#else // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
+        {
+#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_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;
+
+            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, 64, k, GGML_TYPE_F16)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, k, GGML_TYPE_Q4_0)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, k, GGML_TYPE_Q4_1)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, k, GGML_TYPE_Q5_0)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, k, GGML_TYPE_Q5_1)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, k, GGML_TYPE_Q8_0)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, k, GGML_TYPE_BF16)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, k, GGML_TYPE_F32)
+            GGML_ABORT("fatal error");
+        }
+    } else if (n_embd == 128 && n_head == 32) {
+#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
+        if (GGML_CUDA_CC_IS_NVIDIA(cc) && turing_mma_available(cc) && k->type != GGML_TYPE_F32 && k->type != GGML_TYPE_BF16) {
+            // use wmma kernel
+            constexpr int K_VECS_PER_BLOCK = 32;
+            constexpr int WARPS_PER_BLOCK = 8;
+
+            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_wmma, 128, 32, k, GGML_TYPE_F16)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 32, k, GGML_TYPE_Q4_0)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 32, k, GGML_TYPE_Q4_1)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 32, k, GGML_TYPE_Q5_0)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 32, k, GGML_TYPE_Q5_1)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 32, k, GGML_TYPE_Q8_0)
+            GGML_ABORT("fatal error");
+        } else {
+#else // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA)
+        {
+#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_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;
+
+            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, 32, k, GGML_TYPE_F16)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, k, GGML_TYPE_Q4_0)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, k, GGML_TYPE_Q4_1)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, k, GGML_TYPE_Q5_0)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, k, GGML_TYPE_Q5_1)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, k, GGML_TYPE_Q8_0)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, k, GGML_TYPE_BF16)
+            LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, k, GGML_TYPE_F32)
+            GGML_ABORT("fatal error");
+        }
+    } else {
+        GGML_ABORT("fatal error");
+    }
+}
+
+bool ggml_cuda_lightning_indexer_supported(int device, const ggml_tensor * dst) {
+    GGML_UNUSED(device);
+
+    const ggml_tensor * q = dst->src[0];
+    const ggml_tensor * k = dst->src[1];
+    const ggml_tensor * w = dst->src[2]; // weights
+    const ggml_tensor * m = dst->src[3]; // mask
+
+    GGML_TENSOR_LOCALS(int64_t, neq,  q, ne)
+    GGML_TENSOR_LOCALS(size_t,  nbq,  q, nb)
+    GGML_TENSOR_LOCALS(int64_t, nek,  k, ne)
+    GGML_TENSOR_LOCALS(size_t,  nbk,  k, nb)
+    GGML_TENSOR_LOCALS(int64_t, new,  w, ne)
+    GGML_TENSOR_LOCALS(size_t,  nbw,  w, nb)
+    GGML_TENSOR_LOCALS(int64_t, nem,  m, ne)
+    GGML_TENSOR_LOCALS(size_t,  nbm,  m, nb)
+    GGML_TENSOR_LOCALS(int64_t, ne, dst, ne)
+    GGML_TENSOR_LOCALS(size_t,  nb, dst, nb)
+
+    if (neq0 != 128) {
+        return false;
+    }
+
+    if (neq1 != 64 && neq1 != 32) {
+        return false;
+    }
+
+    // alignment checks
+    for (const ggml_tensor * t : {q, k}) {
+        if (ggml_is_quantized(t->type)) {
+            continue;
+        }
+        for (size_t i = 1; i < GGML_MAX_DIMS; ++i) {
+            if (t->nb[i] % 16 != 0) {
+                return false;
+            }
+        }
+    }
+
+    switch(k->type) {
+        case GGML_TYPE_F32:
+        case GGML_TYPE_BF16:
+        case GGML_TYPE_F16:
+        case GGML_TYPE_Q8_0:
+        case GGML_TYPE_Q5_1:
+        case GGML_TYPE_Q5_0:
+        case GGML_TYPE_Q4_1:
+        case GGML_TYPE_Q4_0:
+            return true;
+        default:
+            return false;
+    }
+}
diff --git a/ggml/src/ggml-cuda/lightning-indexer.cuh b/ggml/src/ggml-cuda/lightning-indexer.cuh
new file mode 100644 (file)
index 0000000..f2fc951
--- /dev/null
@@ -0,0 +1,4 @@
+#include "common.cuh"
+
+void ggml_cuda_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor * dst);
+bool ggml_cuda_lightning_indexer_supported(int device, const ggml_tensor * dst);