From 50c49c83a577dcf8df328717dd2eb4d0b3d7909c Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stanis=C5=82aw=20Szymczyk?= Date: Fri, 10 Jul 2026 22:43:16 +0200 Subject: [PATCH 01/12] cuda : CUDA GGML_OP_LIGHTNING_INDEXER implementation (generic vector kernel + wmma kernel) --- ggml/src/ggml-cuda/ggml-cuda.cu | 6 + ggml/src/ggml-cuda/lightning-indexer.cu | 591 +++++++++++++++++++++++ ggml/src/ggml-cuda/lightning-indexer.cuh | 4 + 3 files changed, 601 insertions(+) create mode 100644 ggml/src/ggml-cuda/lightning-indexer.cu create mode 100644 ggml/src/ggml-cuda/lightning-indexer.cuh diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index 98816f885cf6..4585a9fdd2bc 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -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 @@ -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_op_lightning_indexer(ctx, dst); + break; default: return false; } @@ -4970,6 +4974,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 index 000000000000..f23a0f625986 --- /dev/null +++ b/ggml/src/ggml-cuda/lightning-indexer.cu @@ -0,0 +1,591 @@ +#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(__CUDA_ARCH__) && __CUDA_ARCH__ >= GGML_CUDA_CC_AMPERE + +typedef union { + int2 i2; + half2 h2[2]; +} half4; + +#include +namespace wmma = nvcuda::wmma; + +template +static __global__ void lightning_indexer_kernel_wmma( + const float * src0, const char * src1, const float * src2, const half * src3, 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 nb01, size_t nb02, size_t nb03, + size_t nb11, size_t nb12, size_t nb13, + size_t nb21, size_t nb22, size_t nb23, + size_t nb31, size_t nb32, size_t nb33, + int64_t ne33 + ) { + + constexpr int K_VECS_PER_BLOCK = 32; + constexpr int WARPS_PER_BLOCK = 8; + 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 *) src0 + i_batch*nb02 + i_stream*nb03; + const float * w_base = (const float *) ((const char *) src2 + i_batch*nb21 + i_stream*nb23); + + // 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*nb01 + 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 *) src1 + i_kv*nb12 + i_stream*nb13); + *(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(); + #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 *) src1 + i_kv*nb12 + i_stream*nb13); + 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 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 frag_acc; + wmma::fill_fragment(frag_acc, 0.0f); + + // load Q fragment + wmma::fragment 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*nb01 + 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 *) src3 + i_batch*nb31 + (i_stream%ne33)*nb33); + 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(__CUDA_ARCH__) && __CUDA_ARCH__ >= GGML_CUDA_CC_AMPERE + +template +static __global__ void lightning_indexer_kernel_wmma( + const float * src0, const char * src1, const float * src2, const half * src3, 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 nb01, size_t nb02, size_t nb03, + size_t nb11, size_t nb12, size_t nb13, + size_t nb21, size_t nb22, size_t nb23, + size_t nb31, size_t nb32, size_t nb33, + int64_t ne33 + ) { + GGML_UNUSED_VARS(src0, src1, src2, dst, + n_stream, n_batch, n_kv, + nb1, nb2, nb3, + nb01, nb02, nb03, + nb11, nb12, nb13, + nb21, nb22, nb23); + NO_DEVICE_CODE; +} + +#endif // defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= GGML_CUDA_CC_AMPERE +#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) && defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= GGML_CUDA_CC_AMPERE + +// 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 +static __global__ void lightning_indexer_kernel_vec( + const float * src0, const char * src1, const float * src2, const half * src3, 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 nb01, size_t nb02, size_t nb03, + size_t nb11, size_t nb12, size_t nb13, + size_t nb21, size_t nb22, size_t nb23, + size_t nb31, size_t nb32, size_t nb33, + int64_t ne33 + ) { + + constexpr int K_VECS_PER_WARP = 8; + constexpr int WARPS_PER_BLOCK = 8; + 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 * (WARPS_PER_BLOCK * K_VECS_PER_WARP); + const int start_kv = start_kv_block + i_warp * K_VECS_PER_WARP; + + const char * q_base = (const char *) src0 + i_batch*nb02 + i_stream*nb03; + const float * w_base = (const float *) ((const char *) src2 + i_batch*nb21 + i_stream*nb23); + + // 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 *) src1 + i_kv*nb12 + i_stream*nb13); + 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(); + #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 *) src1 + i_kv*nb12 + i_stream*nb13); + 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*nb01 + 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[WARPS_PER_BLOCK * K_VECS_PER_WARP]; + + 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 < WARPS_PER_BLOCK * K_VECS_PER_WARP) { + int i_kv = start_kv_block + tid; + if (i_kv < n_kv) { + const half * m_base = (const half *) ((const char *) src3 + i_batch*nb31 + (i_stream%ne33)*nb33); + 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 DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel, n_embd, n_head, type_K) \ + template __global__ void lightning_indexer_kernel( \ + const float * src0, const char * src1, const float * src2, const half * src3, 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 nb01, size_t nb02, size_t nb03, \ + size_t nb11, size_t nb12, size_t nb13, \ + size_t nb21, size_t nb22, size_t nb23, \ + size_t nb31, size_t nb32, size_t nb33, \ + int64_t ne33); + +#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, GGML_TYPE_F16) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, GGML_TYPE_Q4_0) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, GGML_TYPE_Q4_1) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, GGML_TYPE_Q5_0) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, GGML_TYPE_Q5_1) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, GGML_TYPE_Q8_0) +#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) + +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_F16) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_Q4_0) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_Q4_1) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_Q5_0) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_Q5_1) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_Q8_0) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_BF16) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_F32) + +#define LIGHTNING_INDEXER_CASE(lightning_indexer_kernel, n_embd, n_head, K, type_K) \ + if (K->type == (type_K)) { \ + lightning_indexer_kernel<<>>( \ + src0_d, src1_d, src2_d, src3_d, dst_d, \ + n_stream, n_batch, n_kv, \ + nb1, nb2, nb3, \ + nb01, nb02, nb03, \ + nb11, nb12, nb13, \ + nb21, nb22, nb23, \ + nb31, nb32, nb33, \ + ne33 \ + ); \ + } else + +void ggml_cuda_op_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { + const ggml_tensor * src0 = dst->src[0]; + const ggml_tensor * src1 = dst->src[1]; + const ggml_tensor * src2 = dst->src[2]; + const ggml_tensor * src3 = dst->src[3]; + + GGML_ASSERT(dst->type == GGML_TYPE_F32); + GGML_ASSERT(src0->type == GGML_TYPE_F32); + GGML_ASSERT(src2->type == GGML_TYPE_F32); + GGML_ASSERT(src3->type == GGML_TYPE_F16); + + GGML_TENSOR_TERNARY_OP_LOCALS + GGML_TENSOR_LOCALS(int64_t, ne3, src3, ne) + GGML_TENSOR_LOCALS(size_t, nb3, src3, nb) + + // input tensor rows must be contiguous + GGML_ASSERT(nb00 == ggml_type_size(src0->type)); + GGML_ASSERT(nb10 == ggml_type_size(src1->type)); + GGML_ASSERT(nb20 == ggml_type_size(src2->type)); + GGML_ASSERT(nb30 == ggml_type_size(src3->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 = src0->ne[0]; + const int n_head = src0->ne[1]; + const int n_batch = src0->ne[2]; + const int n_stream = src0->ne[3]; + const int n_kv = src1->ne[2]; + + const float * src0_d = (const float *) src0->data; + const char * src1_d = (const char *) src1->data; + const float * src2_d = (const float *) src2->data; + const half * src3_d = (const half *) src3->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) && ampere_mma_available(cc) && src1->type != GGML_TYPE_F32 && src1->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, src1, GGML_TYPE_F16) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, src1, GGML_TYPE_Q4_0) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, src1, GGML_TYPE_Q4_1) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, src1, GGML_TYPE_Q5_0) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, src1, GGML_TYPE_Q5_1) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, src1, 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, src1, GGML_TYPE_F16) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_Q4_0) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_Q4_1) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_Q5_0) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_Q5_1) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_Q8_0) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_BF16) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, 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) && ampere_mma_available(cc) && src1->type != GGML_TYPE_F32 && src1->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, src1, GGML_TYPE_F16) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 32, src1, GGML_TYPE_Q4_0) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 32, src1, GGML_TYPE_Q4_1) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 32, src1, GGML_TYPE_Q5_0) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 32, src1, GGML_TYPE_Q5_1) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 32, src1, 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, src1, GGML_TYPE_F16) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, src1, GGML_TYPE_Q4_0) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, src1, GGML_TYPE_Q4_1) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, src1, GGML_TYPE_Q5_0) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, src1, GGML_TYPE_Q5_1) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, src1, GGML_TYPE_Q8_0) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, src1, GGML_TYPE_BF16) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, src1, 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 * src0 = dst->src[0]; + const ggml_tensor * src1 = dst->src[1]; + const ggml_tensor * src2 = dst->src[2]; + const ggml_tensor * src3 = dst->src[3]; + + GGML_TENSOR_TERNARY_OP_LOCALS + GGML_TENSOR_LOCALS(int64_t, ne3, src3, ne) + GGML_TENSOR_LOCALS(size_t, nb3, src3, nb) + + if (ne00 != 128) { + return false; + } + + if (ne01 != 64 && ne01 != 32) { + return false; + } + + switch(src1->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 index 000000000000..a9e7527aee3d --- /dev/null +++ b/ggml/src/ggml-cuda/lightning-indexer.cuh @@ -0,0 +1,4 @@ +#include "common.cuh" + +void ggml_cuda_op_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor * dst); +bool ggml_cuda_lightning_indexer_supported(int device, const ggml_tensor * dst); From 1ebcf222cb22c406af30aeb69f408166e5ccaabb Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stanis=C5=82aw=20Szymczyk?= Date: Mon, 13 Jul 2026 20:58:59 +0200 Subject: [PATCH 02/12] chore : remove indentation of #pragma unroll --- ggml/src/ggml-cuda/lightning-indexer.cu | 26 ++++++++++++------------- 1 file changed, 13 insertions(+), 13 deletions(-) diff --git a/ggml/src/ggml-cuda/lightning-indexer.cu b/ggml/src/ggml-cuda/lightning-indexer.cu index f23a0f625986..f10d6decf88d 100644 --- a/ggml/src/ggml-cuda/lightning-indexer.cu +++ b/ggml/src/ggml-cuda/lightning-indexer.cu @@ -59,7 +59,7 @@ static __global__ void lightning_indexer_kernel_wmma( // 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 +#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); @@ -77,7 +77,7 @@ static __global__ void lightning_indexer_kernel_wmma( constexpr int n_k = K_VECS_PER_BLOCK * (n_embd / 4); if constexpr (type_K == GGML_TYPE_F16) { - #pragma unroll +#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); @@ -91,7 +91,7 @@ static __global__ void lightning_indexer_kernel_wmma( } } else { constexpr dequantize_V_t dequantize_k = get_dequantize_V(); - #pragma unroll +#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); @@ -132,7 +132,7 @@ static __global__ void lightning_indexer_kernel_wmma( float4 q_next[n_q_next]; if (i_head_next < n_head) { - #pragma unroll +#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); @@ -149,7 +149,7 @@ static __global__ void lightning_indexer_kernel_wmma( // write preloaded Q tile to shared memory if (i_head_next < n_head) { - #pragma unroll +#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); @@ -169,7 +169,7 @@ static __global__ void lightning_indexer_kernel_wmma( const float w_val = w_shared[i_head_0 + h]; float sum = 0.0f; - #pragma unroll +#pragma unroll for (int w = 0; w < WARPS_PER_BLOCK; ++w) { sum += qk_shared[w][h][k]; } @@ -189,7 +189,7 @@ static __global__ void lightning_indexer_kernel_wmma( // accumulate result over heads if (tid < K_VECS_PER_BLOCK) { - #pragma unroll +#pragma unroll for (int i_head = 0; i_head < HEADS_PER_INNER_LOOP; ++i_head) { score_k += qk_shared[0][i_head][tid]; } @@ -275,7 +275,7 @@ static __global__ void lightning_indexer_kernel_vec( if constexpr (type_K == GGML_TYPE_F32) { // direct copy of float4 - #pragma unroll +#pragma unroll for (int k = 0; k < K_VECS_PER_WARP; ++k) { int i_kv = start_kv + k; if (i_kv < n_kv) { @@ -288,7 +288,7 @@ static __global__ void lightning_indexer_kernel_vec( } else { // dequantize remaining types to float constexpr dequantize_V_t dequantize_k = get_dequantize_V(); - #pragma unroll +#pragma unroll for (int k = 0; k < K_VECS_PER_WARP; ++k) { int i_kv = start_kv + k; if (i_kv < n_kv) { @@ -316,7 +316,7 @@ static __global__ void lightning_indexer_kernel_vec( } constexpr int n_q = n_head_inner * (n_embd / 4); - #pragma unroll +#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; @@ -335,7 +335,7 @@ static __global__ void lightning_indexer_kernel_vec( // dot product of floats const float4 q_vec = q_shared_f[i_head_inner][i_lane]; - #pragma unroll +#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); @@ -343,7 +343,7 @@ static __global__ void lightning_indexer_kernel_vec( ggml_cuda_mad(qk[k], q_vec.w, k_reg_f[k].w); } - #pragma unroll +#pragma unroll for (int k = 0; k < K_VECS_PER_WARP; ++k) { float sum = warp_reduce_sum(qk[k]); @@ -363,7 +363,7 @@ static __global__ void lightning_indexer_kernel_vec( __shared__ float dst_shared[WARPS_PER_BLOCK * K_VECS_PER_WARP]; if (i_lane == 0) { - #pragma unroll +#pragma unroll for (int k = 0; k < K_VECS_PER_WARP; ++k) { dst_shared[i_warp * K_VECS_PER_WARP + k] = score_k[k]; } From 07bfc3492010269ca2d190982f63602cd81fe716 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stanis=C5=82aw=20Szymczyk?= Date: Tue, 14 Jul 2026 10:53:27 +0200 Subject: [PATCH 03/12] cuda : remove unnecessary kernel template declarations --- ggml/src/ggml-cuda/lightning-indexer.cu | 29 ------------------------- 1 file changed, 29 deletions(-) diff --git a/ggml/src/ggml-cuda/lightning-indexer.cu b/ggml/src/ggml-cuda/lightning-indexer.cu index f10d6decf88d..27f2cb334b19 100644 --- a/ggml/src/ggml-cuda/lightning-indexer.cu +++ b/ggml/src/ggml-cuda/lightning-indexer.cu @@ -383,35 +383,6 @@ static __global__ void lightning_indexer_kernel_vec( } } -#define DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel, n_embd, n_head, type_K) \ - template __global__ void lightning_indexer_kernel( \ - const float * src0, const char * src1, const float * src2, const half * src3, 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 nb01, size_t nb02, size_t nb03, \ - size_t nb11, size_t nb12, size_t nb13, \ - size_t nb21, size_t nb22, size_t nb23, \ - size_t nb31, size_t nb32, size_t nb33, \ - int64_t ne33); - -#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) -DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, GGML_TYPE_F16) -DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, GGML_TYPE_Q4_0) -DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, GGML_TYPE_Q4_1) -DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, GGML_TYPE_Q5_0) -DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, GGML_TYPE_Q5_1) -DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, GGML_TYPE_Q8_0) -#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) - -DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_F16) -DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_Q4_0) -DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_Q4_1) -DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_Q5_0) -DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_Q5_1) -DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_Q8_0) -DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_BF16) -DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_F32) - #define LIGHTNING_INDEXER_CASE(lightning_indexer_kernel, n_embd, n_head, K, type_K) \ if (K->type == (type_K)) { \ lightning_indexer_kernel<<>>( \ From 8d59c209463ec7b43e153a53c3e3ae568e3b567e Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stanis=C5=82aw=20Szymczyk?= Date: Tue, 14 Jul 2026 11:48:16 +0200 Subject: [PATCH 04/12] cuda : add WARPS_PER_BLOCK and K_VECS_PER_BLOCK template parameters in lightning indexer kernels to avoid duplication of constants. --- ggml/src/ggml-cuda/lightning-indexer.cu | 20 +++++++++----------- 1 file changed, 9 insertions(+), 11 deletions(-) diff --git a/ggml/src/ggml-cuda/lightning-indexer.cu b/ggml/src/ggml-cuda/lightning-indexer.cu index 27f2cb334b19..0cd2ea1933a4 100644 --- a/ggml/src/ggml-cuda/lightning-indexer.cu +++ b/ggml/src/ggml-cuda/lightning-indexer.cu @@ -14,7 +14,7 @@ typedef union { #include namespace wmma = nvcuda::wmma; -template +template static __global__ void lightning_indexer_kernel_wmma( const float * src0, const char * src1, const float * src2, const half * src3, float * dst, int64_t n_stream, int64_t n_batch, int64_t n_kv, @@ -26,8 +26,6 @@ static __global__ void lightning_indexer_kernel_wmma( int64_t ne33 ) { - constexpr int K_VECS_PER_BLOCK = 32; - constexpr int WARPS_PER_BLOCK = 8; 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; @@ -213,7 +211,7 @@ static __global__ void lightning_indexer_kernel_wmma( #else // defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= GGML_CUDA_CC_AMPERE -template +template static __global__ void lightning_indexer_kernel_wmma( const float * src0, const char * src1, const float * src2, const half * src3, float * dst, int64_t n_stream, int64_t n_batch, int64_t n_kv, @@ -240,7 +238,7 @@ static __global__ void lightning_indexer_kernel_wmma( // thanks to that one warp operating on float4 processes whole indexer K/Q vectors // 32 * 4 = 128 (n_embd) -template +template static __global__ void lightning_indexer_kernel_vec( const float * src0, const char * src1, const float * src2, const half * src3, float * dst, int64_t n_stream, int64_t n_batch, int64_t n_kv, @@ -252,8 +250,7 @@ static __global__ void lightning_indexer_kernel_vec( int64_t ne33 ) { - constexpr int K_VECS_PER_WARP = 8; - constexpr int WARPS_PER_BLOCK = 8; + 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; @@ -263,7 +260,7 @@ static __global__ void lightning_indexer_kernel_vec( 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 * (WARPS_PER_BLOCK * K_VECS_PER_WARP); + 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 *) src0 + i_batch*nb02 + i_stream*nb03; @@ -360,7 +357,7 @@ static __global__ void lightning_indexer_kernel_vec( // phase 4 - store outputs to shared memory - __shared__ float dst_shared[WARPS_PER_BLOCK * K_VECS_PER_WARP]; + __shared__ float dst_shared[K_VECS_PER_BLOCK]; if (i_lane == 0) { #pragma unroll @@ -373,7 +370,7 @@ static __global__ void lightning_indexer_kernel_vec( // phase 5 - write from shared memory to VRAM in coalesced manner - if (tid < WARPS_PER_BLOCK * K_VECS_PER_WARP) { + 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 *) src3 + i_batch*nb31 + (i_stream%ne33)*nb33); @@ -385,7 +382,8 @@ static __global__ void lightning_indexer_kernel_vec( #define LIGHTNING_INDEXER_CASE(lightning_indexer_kernel, n_embd, n_head, K, type_K) \ if (K->type == (type_K)) { \ - lightning_indexer_kernel<<>>( \ + lightning_indexer_kernel \ + <<>>( \ src0_d, src1_d, src2_d, src3_d, dst_d, \ n_stream, n_batch, n_kv, \ nb1, nb2, nb3, \ From d372562b7bd23c6e10fa453798a3bfc1e9c44ef0 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stanis=C5=82aw=20Szymczyk?= Date: Tue, 14 Jul 2026 11:43:01 +0000 Subject: [PATCH 05/12] cuda : relax MMA architecture requirements to Turing in lightning indexer implementation --- ggml/src/ggml-cuda/lightning-indexer.cu | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/ggml/src/ggml-cuda/lightning-indexer.cu b/ggml/src/ggml-cuda/lightning-indexer.cu index 0cd2ea1933a4..355d78470fb1 100644 --- a/ggml/src/ggml-cuda/lightning-indexer.cu +++ b/ggml/src/ggml-cuda/lightning-indexer.cu @@ -4,7 +4,7 @@ #include "convert.cuh" #if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) -#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= GGML_CUDA_CC_AMPERE +#if defined(TURING_MMA_AVAILABLE) typedef union { int2 i2; @@ -209,7 +209,7 @@ static __global__ void lightning_indexer_kernel_wmma( } } -#else // defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= GGML_CUDA_CC_AMPERE +#else // defined(TURING_MMA_AVAILABLE) template static __global__ void lightning_indexer_kernel_wmma( @@ -231,8 +231,8 @@ static __global__ void lightning_indexer_kernel_wmma( NO_DEVICE_CODE; } -#endif // defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= GGML_CUDA_CC_AMPERE -#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) && defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= GGML_CUDA_CC_AMPERE +#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 @@ -439,7 +439,7 @@ void ggml_cuda_op_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor if (n_embd == 128 && n_head == 64) { #if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) - if (GGML_CUDA_CC_IS_NVIDIA(cc) && ampere_mma_available(cc) && src1->type != GGML_TYPE_F32 && src1->type != GGML_TYPE_BF16) { + if (GGML_CUDA_CC_IS_NVIDIA(cc) && turing_mma_available(cc) && src1->type != GGML_TYPE_F32 && src1->type != GGML_TYPE_BF16) { // use wmma kernel constexpr int K_VECS_PER_BLOCK = 32; constexpr int WARPS_PER_BLOCK = 8; @@ -480,7 +480,7 @@ void ggml_cuda_op_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor } } else if (n_embd == 128 && n_head == 32) { #if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) - if (GGML_CUDA_CC_IS_NVIDIA(cc) && ampere_mma_available(cc) && src1->type != GGML_TYPE_F32 && src1->type != GGML_TYPE_BF16) { + if (GGML_CUDA_CC_IS_NVIDIA(cc) && turing_mma_available(cc) && src1->type != GGML_TYPE_F32 && src1->type != GGML_TYPE_BF16) { // use wmma kernel constexpr int K_VECS_PER_BLOCK = 32; constexpr int WARPS_PER_BLOCK = 8; From 9937427cd9c8d03341417a9a0ee53660628a3d7f Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stanis=C5=82aw=20Szymczyk?= Date: Tue, 14 Jul 2026 14:33:00 +0200 Subject: [PATCH 06/12] chore : renamed variables --- ggml/src/ggml-cuda/lightning-indexer.cu | 231 +++++++++++++----------- 1 file changed, 123 insertions(+), 108 deletions(-) diff --git a/ggml/src/ggml-cuda/lightning-indexer.cu b/ggml/src/ggml-cuda/lightning-indexer.cu index 355d78470fb1..39cbc940f14a 100644 --- a/ggml/src/ggml-cuda/lightning-indexer.cu +++ b/ggml/src/ggml-cuda/lightning-indexer.cu @@ -16,14 +16,14 @@ namespace wmma = nvcuda::wmma; template static __global__ void lightning_indexer_kernel_wmma( - const float * src0, const char * src1, const float * src2, const half * src3, float * dst, + 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 nb01, size_t nb02, size_t nb03, - size_t nb11, size_t nb12, size_t nb13, - size_t nb21, size_t nb22, size_t nb23, - size_t nb31, size_t nb32, size_t nb33, - int64_t ne33 + 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; @@ -40,8 +40,8 @@ static __global__ void lightning_indexer_kernel_wmma( // 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 *) src0 + i_batch*nb02 + i_stream*nb03; - const float * w_base = (const float *) ((const char *) src2 + i_batch*nb21 + i_stream*nb23); + 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 @@ -61,7 +61,7 @@ static __global__ void lightning_indexer_kernel_wmma( 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*nb01 + i_embd*sizeof(float4)); + 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)); @@ -81,7 +81,7 @@ static __global__ void lightning_indexer_kernel_wmma( 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 *) src1 + i_kv*nb12 + i_stream*nb13); + 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); @@ -95,7 +95,7 @@ static __global__ void lightning_indexer_kernel_wmma( 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 *) src1 + i_kv*nb12 + i_stream*nb13); + 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); @@ -134,7 +134,7 @@ static __global__ void lightning_indexer_kernel_wmma( 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*nb01 + i_embd*sizeof(float4)); + q_next[i_q_next++] = *(const float4 *) (q_base + i_head*nbq1 + i_embd*sizeof(float4)); } } @@ -202,7 +202,7 @@ static __global__ void lightning_indexer_kernel_wmma( 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 *) src3 + i_batch*nb31 + (i_stream%ne33)*nb33); + 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]); } @@ -213,21 +213,22 @@ static __global__ void lightning_indexer_kernel_wmma( template static __global__ void lightning_indexer_kernel_wmma( - const float * src0, const char * src1, const float * src2, const half * src3, float * dst, + 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 nb01, size_t nb02, size_t nb03, - size_t nb11, size_t nb12, size_t nb13, - size_t nb21, size_t nb22, size_t nb23, - size_t nb31, size_t nb32, size_t nb33, - int64_t ne33 + 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(src0, src1, src2, dst, + GGML_UNUSED_VARS(q, k, w, m, dst, n_stream, n_batch, n_kv, nb1, nb2, nb3, - nb01, nb02, nb03, - nb11, nb12, nb13, - nb21, nb22, nb23); + nbq1, nbq2, nbq3, + nbk1, nbk2, nbk3, + nbw1, nbw2, nbw3, + nem3); NO_DEVICE_CODE; } @@ -240,14 +241,14 @@ static __global__ void lightning_indexer_kernel_wmma( template static __global__ void lightning_indexer_kernel_vec( - const float * src0, const char * src1, const float * src2, const half * src3, float * dst, + 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 nb01, size_t nb02, size_t nb03, - size_t nb11, size_t nb12, size_t nb13, - size_t nb21, size_t nb22, size_t nb23, - size_t nb31, size_t nb32, size_t nb33, - int64_t ne33 + 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; @@ -263,8 +264,8 @@ static __global__ void lightning_indexer_kernel_vec( 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 *) src0 + i_batch*nb02 + i_stream*nb03; - const float * w_base = (const float *) ((const char *) src2 + i_batch*nb21 + i_stream*nb23); + 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 @@ -276,7 +277,7 @@ static __global__ void lightning_indexer_kernel_vec( 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 *) src1 + i_kv*nb12 + i_stream*nb13); + 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); @@ -289,7 +290,7 @@ static __global__ void lightning_indexer_kernel_vec( 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 *) src1 + i_kv*nb12 + i_stream*nb13); + 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); @@ -318,7 +319,7 @@ static __global__ void lightning_indexer_kernel_vec( 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*nb01 + i_embd*sizeof(float4)); + q_shared_f[i_head_inner][i_embd] = *(const float4 *) (q_base + i_head*nbq1 + i_embd*sizeof(float4)); } __syncthreads(); @@ -373,7 +374,7 @@ static __global__ void lightning_indexer_kernel_vec( 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 *) src3 + i_batch*nb31 + (i_stream%ne33)*nb33); + 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]); } @@ -384,37 +385,44 @@ static __global__ void lightning_indexer_kernel_vec( if (K->type == (type_K)) { \ lightning_indexer_kernel \ <<>>( \ - src0_d, src1_d, src2_d, src3_d, dst_d, \ + q_d, k_d, w_d, m_d, dst_d, \ n_stream, n_batch, n_kv, \ nb1, nb2, nb3, \ - nb01, nb02, nb03, \ - nb11, nb12, nb13, \ - nb21, nb22, nb23, \ - nb31, nb32, nb33, \ - ne33 \ + nbq1, nbq2, nbq3, \ + nbk1, nbk2, nbk3, \ + nbw1, nbw2, nbw3, \ + nbm1, nbm2, nbm3, \ + nem3 \ ); \ } else void ggml_cuda_op_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { - const ggml_tensor * src0 = dst->src[0]; - const ggml_tensor * src1 = dst->src[1]; - const ggml_tensor * src2 = dst->src[2]; - const ggml_tensor * src3 = dst->src[3]; + 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(src0->type == GGML_TYPE_F32); - GGML_ASSERT(src2->type == GGML_TYPE_F32); - GGML_ASSERT(src3->type == GGML_TYPE_F16); - - GGML_TENSOR_TERNARY_OP_LOCALS - GGML_TENSOR_LOCALS(int64_t, ne3, src3, ne) - GGML_TENSOR_LOCALS(size_t, nb3, src3, nb) + 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(nb00 == ggml_type_size(src0->type)); - GGML_ASSERT(nb10 == ggml_type_size(src1->type)); - GGML_ASSERT(nb20 == ggml_type_size(src2->type)); - GGML_ASSERT(nb30 == ggml_type_size(src3->type)); + 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)); @@ -422,24 +430,24 @@ void ggml_cuda_op_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor GGML_ASSERT(nb1 <= nb2); GGML_ASSERT(nb2 <= nb3); - const int n_embd = src0->ne[0]; - const int n_head = src0->ne[1]; - const int n_batch = src0->ne[2]; - const int n_stream = src0->ne[3]; - const int n_kv = src1->ne[2]; + 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 * src0_d = (const float *) src0->data; - const char * src1_d = (const char *) src1->data; - const float * src2_d = (const float *) src2->data; - const half * src3_d = (const half *) src3->data; - float * dst_d = (float *) dst->data; + 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) && src1->type != GGML_TYPE_F32 && src1->type != GGML_TYPE_BF16) { + 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; @@ -448,12 +456,12 @@ void ggml_cuda_op_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor 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, src1, GGML_TYPE_F16) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, src1, GGML_TYPE_Q4_0) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, src1, GGML_TYPE_Q4_1) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, src1, GGML_TYPE_Q5_0) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, src1, GGML_TYPE_Q5_1) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, src1, GGML_TYPE_Q8_0) + 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) @@ -468,19 +476,19 @@ void ggml_cuda_op_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor 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, src1, GGML_TYPE_F16) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_Q4_0) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_Q4_1) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_Q5_0) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_Q5_1) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_Q8_0) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_BF16) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_F32) + 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) && src1->type != GGML_TYPE_F32 && src1->type != GGML_TYPE_BF16) { + 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; @@ -489,12 +497,12 @@ void ggml_cuda_op_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor 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, src1, GGML_TYPE_F16) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 32, src1, GGML_TYPE_Q4_0) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 32, src1, GGML_TYPE_Q4_1) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 32, src1, GGML_TYPE_Q5_0) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 32, src1, GGML_TYPE_Q5_1) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 32, src1, GGML_TYPE_Q8_0) + 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) @@ -509,14 +517,14 @@ void ggml_cuda_op_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor 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, src1, GGML_TYPE_F16) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, src1, GGML_TYPE_Q4_0) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, src1, GGML_TYPE_Q4_1) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, src1, GGML_TYPE_Q5_0) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, src1, GGML_TYPE_Q5_1) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, src1, GGML_TYPE_Q8_0) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, src1, GGML_TYPE_BF16) - LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 32, src1, GGML_TYPE_F32) + 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 { @@ -527,24 +535,31 @@ void ggml_cuda_op_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor bool ggml_cuda_lightning_indexer_supported(int device, const ggml_tensor * dst) { GGML_UNUSED(device); - const ggml_tensor * src0 = dst->src[0]; - const ggml_tensor * src1 = dst->src[1]; - const ggml_tensor * src2 = dst->src[2]; - const ggml_tensor * src3 = dst->src[3]; - - GGML_TENSOR_TERNARY_OP_LOCALS - GGML_TENSOR_LOCALS(int64_t, ne3, src3, ne) - GGML_TENSOR_LOCALS(size_t, nb3, src3, nb) - - if (ne00 != 128) { + 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 (ne01 != 64 && ne01 != 32) { + if (neq1 != 64 && neq1 != 32) { return false; } - switch(src1->type) { + switch(k->type) { case GGML_TYPE_F32: case GGML_TYPE_BF16: case GGML_TYPE_F16: From 44b29b7897376f4af35558b6eb7aefe3f51b35b1 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stanis=C5=82aw=20Szymczyk?= Date: Tue, 14 Jul 2026 14:43:48 +0200 Subject: [PATCH 07/12] chore : rename ggml_cuda_op_lightning_indexer() to ggml_cuda_lightning_indexer() --- ggml/src/ggml-cuda/ggml-cuda.cu | 2 +- ggml/src/ggml-cuda/lightning-indexer.cu | 2 +- ggml/src/ggml-cuda/lightning-indexer.cuh | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index 4585a9fdd2bc..58139f621309 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -2259,7 +2259,7 @@ static bool ggml_cuda_compute_forward(ggml_backend_cuda_context & ctx, struct gg ggml_cuda_op_fill(ctx, dst); break; case GGML_OP_LIGHTNING_INDEXER: - ggml_cuda_op_lightning_indexer(ctx, dst); + ggml_cuda_lightning_indexer(ctx, dst); break; default: return false; diff --git a/ggml/src/ggml-cuda/lightning-indexer.cu b/ggml/src/ggml-cuda/lightning-indexer.cu index 39cbc940f14a..ef5196cebcbc 100644 --- a/ggml/src/ggml-cuda/lightning-indexer.cu +++ b/ggml/src/ggml-cuda/lightning-indexer.cu @@ -396,7 +396,7 @@ static __global__ void lightning_indexer_kernel_vec( ); \ } else -void ggml_cuda_op_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { +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 diff --git a/ggml/src/ggml-cuda/lightning-indexer.cuh b/ggml/src/ggml-cuda/lightning-indexer.cuh index a9e7527aee3d..f2fc95181339 100644 --- a/ggml/src/ggml-cuda/lightning-indexer.cuh +++ b/ggml/src/ggml-cuda/lightning-indexer.cuh @@ -1,4 +1,4 @@ #include "common.cuh" -void ggml_cuda_op_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor * dst); +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); From 8e683483e36b39cb52c1f6ece5ff671ae2548051 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stanis=C5=82aw=20Szymczyk?= Date: Tue, 14 Jul 2026 14:49:44 +0200 Subject: [PATCH 08/12] chore : TODO for AMD rocWMMA --- ggml/src/ggml-cuda/lightning-indexer.cu | 1 + 1 file changed, 1 insertion(+) diff --git a/ggml/src/ggml-cuda/lightning-indexer.cu b/ggml/src/ggml-cuda/lightning-indexer.cu index ef5196cebcbc..2ccf839602d4 100644 --- a/ggml/src/ggml-cuda/lightning-indexer.cu +++ b/ggml/src/ggml-cuda/lightning-indexer.cu @@ -11,6 +11,7 @@ typedef union { half2 h2[2]; } half4; +// TODO add support for AMD cards via rocWMMA #include namespace wmma = nvcuda::wmma; From aaa8968fe763e939b9249ce584a1ed33aff5d760 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stanis=C5=82aw=20Szymczyk?= Date: Tue, 14 Jul 2026 16:19:59 +0200 Subject: [PATCH 09/12] chore : whitespace formatting --- ggml/src/ggml-cuda/lightning-indexer.cu | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/ggml/src/ggml-cuda/lightning-indexer.cu b/ggml/src/ggml-cuda/lightning-indexer.cu index 2ccf839602d4..7041f85c2867 100644 --- a/ggml/src/ggml-cuda/lightning-indexer.cu +++ b/ggml/src/ggml-cuda/lightning-indexer.cu @@ -403,10 +403,10 @@ void ggml_cuda_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor * 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_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) From 01878d939b78993c72d0afe6ca636d4c700d2f3d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stanis=C5=82aw=20Szymczyk?= Date: Tue, 14 Jul 2026 17:56:08 +0200 Subject: [PATCH 10/12] chore : another variable rename to fix problems caused by shadowing --- ggml/src/ggml-cuda/lightning-indexer.cu | 28 ++++++++++++------------- 1 file changed, 14 insertions(+), 14 deletions(-) diff --git a/ggml/src/ggml-cuda/lightning-indexer.cu b/ggml/src/ggml-cuda/lightning-indexer.cu index 7041f85c2867..9570c082bd07 100644 --- a/ggml/src/ggml-cuda/lightning-indexer.cu +++ b/ggml/src/ggml-cuda/lightning-indexer.cu @@ -17,7 +17,7 @@ namespace wmma = nvcuda::wmma; template static __global__ void lightning_indexer_kernel_wmma( - const float * q, const char * k, const float * w, const half * m, float * dst, + 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, @@ -41,8 +41,8 @@ static __global__ void lightning_indexer_kernel_wmma( // 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); + 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 @@ -82,7 +82,7 @@ static __global__ void lightning_indexer_kernel_wmma( 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); + 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); @@ -96,7 +96,7 @@ static __global__ void lightning_indexer_kernel_wmma( 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); + 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); @@ -203,7 +203,7 @@ static __global__ void lightning_indexer_kernel_wmma( 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); + 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]); } @@ -214,7 +214,7 @@ static __global__ void lightning_indexer_kernel_wmma( template static __global__ void lightning_indexer_kernel_wmma( - const float * q, const char * k, const float * w, const half * m, float * dst, + 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, @@ -223,7 +223,7 @@ static __global__ void lightning_indexer_kernel_wmma( size_t nbm1, size_t nbm2, size_t nbm3, int64_t nem3 ) { - GGML_UNUSED_VARS(q, k, w, m, dst, + GGML_UNUSED_VARS(Q, K, W, M, dst, n_stream, n_batch, n_kv, nb1, nb2, nb3, nbq1, nbq2, nbq3, @@ -242,7 +242,7 @@ static __global__ void lightning_indexer_kernel_wmma( template static __global__ void lightning_indexer_kernel_vec( - const float * q, const char * k, const float * w, const half * m, float * dst, + 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, @@ -265,8 +265,8 @@ static __global__ void lightning_indexer_kernel_vec( 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); + 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 @@ -278,7 +278,7 @@ static __global__ void lightning_indexer_kernel_vec( 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); + 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); @@ -291,7 +291,7 @@ static __global__ void lightning_indexer_kernel_vec( 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); + 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); @@ -375,7 +375,7 @@ static __global__ void lightning_indexer_kernel_vec( 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); + 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]); } From 3ab3bb731d32b521236245cbfe7ed6e8580a760a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stanis=C5=82aw=20Szymczyk?= Date: Tue, 14 Jul 2026 19:02:01 +0200 Subject: [PATCH 11/12] chore : yet another rename, this time uppercased all constants --- ggml/src/ggml-cuda/lightning-indexer.cu | 92 ++++++++++++------------- 1 file changed, 46 insertions(+), 46 deletions(-) diff --git a/ggml/src/ggml-cuda/lightning-indexer.cu b/ggml/src/ggml-cuda/lightning-indexer.cu index 9570c082bd07..3596b0a1a668 100644 --- a/ggml/src/ggml-cuda/lightning-indexer.cu +++ b/ggml/src/ggml-cuda/lightning-indexer.cu @@ -15,7 +15,7 @@ typedef union { #include namespace wmma = nvcuda::wmma; -template +template 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, @@ -30,7 +30,7 @@ static __global__ void lightning_indexer_kernel_wmma( 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; + constexpr int N_EMBD_PADDED = N_EMBD + 8; const int i_batch = blockIdx.y; const int i_stream = blockIdx.z; @@ -46,22 +46,22 @@ static __global__ void lightning_indexer_kernel_wmma( // 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]; + __shared__ float w_shared[N_HEAD]; + __shared__ int2 q_shared_h[HEADS_PER_INNER_LOOP][N_EMBD_PADDED / 4]; - if (tid < n_head) { + 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); + // 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; + 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); + 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)); @@ -71,15 +71,15 @@ static __global__ void lightning_indexer_kernel_wmma( // 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]; + __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); + constexpr int n_k = K_VECS_PER_BLOCK * (N_EMBD / 4); - if constexpr (type_K == GGML_TYPE_F16) { + 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_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); @@ -89,11 +89,11 @@ static __global__ void lightning_indexer_kernel_wmma( } } } else { - constexpr dequantize_V_t dequantize_k = get_dequantize_V(); + constexpr dequantize_V_t dequantize_k = get_dequantize_V(); #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_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); @@ -112,11 +112,11 @@ static __global__ void lightning_indexer_kernel_wmma( // load K fragment wmma::fragment frag_k; - wmma::load_matrix_sync(frag_k, (half*) &k_shared_h[0][i_warp * K_EMBD_PER_INNER_LOOP / 4], n_embd_padded); + 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) { + 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 @@ -125,16 +125,16 @@ static __global__ void lightning_indexer_kernel_wmma( // load Q fragment wmma::fragment frag_q; - wmma::load_matrix_sync(frag_q, (half*) &q_shared_h[0][i_warp * K_EMBD_PER_INNER_LOOP / 4], n_embd_padded); + 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]; + float4 q_next[N_Q_NEXT]; - if (i_head_next < n_head) { + 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); + 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)); } } @@ -147,11 +147,11 @@ static __global__ void lightning_indexer_kernel_wmma( __syncthreads(); // write preloaded Q tile to shared memory - if (i_head_next < n_head) { + 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); + 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)); @@ -212,7 +212,7 @@ static __global__ void lightning_indexer_kernel_wmma( #else // defined(TURING_MMA_AVAILABLE) -template +template 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, @@ -238,9 +238,9 @@ static __global__ void lightning_indexer_kernel_wmma( // 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) +// 32 * 4 = 128 (N_EMBD) -template +template 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, @@ -272,7 +272,7 @@ static __global__ void lightning_indexer_kernel_vec( float4 k_reg_f[K_VECS_PER_WARP]; - if constexpr (type_K == GGML_TYPE_F32) { + if constexpr (TYPE_K == GGML_TYPE_F32) { // direct copy of float4 #pragma unroll for (int k = 0; k < K_VECS_PER_WARP; ++k) { @@ -286,7 +286,7 @@ static __global__ void lightning_indexer_kernel_vec( } } else { // dequantize remaining types to float - constexpr dequantize_V_t dequantize_k = get_dequantize_V(); + constexpr dequantize_V_t dequantize_k = get_dequantize_V(); #pragma unroll for (int k = 0; k < K_VECS_PER_WARP; ++k) { int i_kv = start_kv + k; @@ -301,25 +301,25 @@ static __global__ void lightning_indexer_kernel_vec( 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; + // 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) { + 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]; + __shared__ float w_shared[N_HEAD_INNER]; + __shared__ float4 q_shared_f[N_HEAD_INNER][N_EMBD / 4]; - if (tid < n_head_inner) { + if (tid < N_HEAD_INNER) { w_shared[tid] = w_base[i_head_0 + tid]; } - constexpr int n_q = n_head_inner * (n_embd / 4); + 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_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); + 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)); } @@ -327,7 +327,7 @@ static __global__ void lightning_indexer_kernel_vec( // phase 3 - calculate lightning indexer scores - for (int i_head_inner = 0; i_head_inner < n_head_inner; ++i_head_inner) { + 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 }; From 37b6a65d5b8f74df8c2f5961987c0db787460fe0 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stanis=C5=82aw=20Szymczyk?= Date: Tue, 14 Jul 2026 19:03:52 +0200 Subject: [PATCH 12/12] cuda : added alignment checks for Q and K tensors in lightning indexer implementation --- ggml/src/ggml-cuda/lightning-indexer.cu | 12 ++++++++++++ 1 file changed, 12 insertions(+) diff --git a/ggml/src/ggml-cuda/lightning-indexer.cu b/ggml/src/ggml-cuda/lightning-indexer.cu index 3596b0a1a668..5edc967e0e92 100644 --- a/ggml/src/ggml-cuda/lightning-indexer.cu +++ b/ggml/src/ggml-cuda/lightning-indexer.cu @@ -560,6 +560,18 @@ bool ggml_cuda_lightning_indexer_supported(int device, const ggml_tensor * dst) 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: