From ecf3c34b0751d4d57c2dc0486f386fccb78d2d62 Mon Sep 17 00:00:00 2001 From: Thiago Padilha Date: Sat, 18 Jul 2026 10:45:21 -0300 Subject: [PATCH 1/4] metal: implement F16 Lightning Indexer MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Implement GGML_OP_LIGHTNING_INDEXER for 128-dimensional, 64-head inputs with F32 queries and weights plus F16 keys and masks. - Add tiled and tail kernels and test KV lengths around 8- and 64-element boundaries. llama-bench (--mmap 1, -fa 1, -p 512, -n 128; d=0/10k/20k/30k): Before: - pp512: 153.73 ± 0.87 t/s - tg128: 8.91 ± 0.04 t/s - pp512 @ d10000: 73.90 ± 0.39 t/s - tg128 @ d10000: 8.66 ± 0.03 t/s - pp512 @ d20000: 45.83 ± 0.18 t/s - tg128 @ d20000: 8.26 ± 0.03 t/s - pp512 @ d30000: 33.40 ± 0.21 t/s - tg128 @ d30000: 7.94 ± 0.01 t/s After: - pp512: 155.19 ± 0.91 t/s - tg128: 8.95 ± 0.04 t/s - pp512 @ d10000: 86.95 ± 0.69 t/s - tg128 @ d10000: 9.00 ± 0.05 t/s - pp512 @ d20000: 62.01 ± 0.45 t/s - tg128 @ d20000: 8.68 ± 0.04 t/s - pp512 @ d30000: 49.18 ± 0.33 t/s - tg128 @ d30000: 8.60 ± 0.02 t/s Assisted-by: Codex --- ggml/src/ggml-metal/ggml-metal-device.cpp | 11 ++ ggml/src/ggml-metal/ggml-metal-device.h | 1 + ggml/src/ggml-metal/ggml-metal-device.m | 13 ++ ggml/src/ggml-metal/ggml-metal-impl.h | 18 +++ ggml/src/ggml-metal/ggml-metal-ops.cpp | 87 +++++++++++++ ggml/src/ggml-metal/ggml-metal-ops.h | 1 + ggml/src/ggml-metal/ggml-metal.metal | 141 ++++++++++++++++++++++ tests/test-backend-ops.cpp | 4 + 8 files changed, 276 insertions(+) diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp index 270c1411a05..e973e0b305f 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.cpp +++ b/ggml/src/ggml-metal/ggml-metal-device.cpp @@ -477,6 +477,17 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_soft_max(ggml_me return res; } +ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning_indexer(ggml_metal_library_t lib, bool tail) { + const char * name = tail ? "kernel_lightning_indexer_f16_tail" : "kernel_lightning_indexer_f16"; + + ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name); + if (!res.pipeline) { + res = ggml_metal_library_compile_pipeline(lib, name, name, nullptr); + } + + return res; +} + ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv(ggml_metal_library_t lib, const ggml_tensor * op) { GGML_ASSERT(op->src[0]->type == GGML_TYPE_F32); GGML_ASSERT(op->src[1]->type == GGML_TYPE_F32); diff --git a/ggml/src/ggml-metal/ggml-metal-device.h b/ggml/src/ggml-metal/ggml-metal-device.h index b36fa8110b5..47872539af2 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.h +++ b/ggml/src/ggml-metal/ggml-metal-device.h @@ -124,6 +124,7 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_cumsum_bl struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_cumsum_add (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_tri (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_soft_max (ggml_metal_library_t lib, const struct ggml_tensor * op); +struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning_indexer (ggml_metal_library_t lib, bool tail); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv_batched (ggml_metal_library_t lib, const struct ggml_tensor * op, int ssm_conv_bs); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan (ggml_metal_library_t lib, const struct ggml_tensor * op); diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index 80e47f2c2f8..eb7d7eae1dd 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -1264,6 +1264,19 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te return false; } return has_simdgroup_mm; // TODO: over-restricted for vec-kernels + case GGML_OP_LIGHTNING_INDEXER: + return has_simdgroup_mm && + op->src[0]->type == GGML_TYPE_F32 && + op->src[1]->type == GGML_TYPE_F16 && + op->src[2]->type == GGML_TYPE_F32 && + op->src[3]->type == GGML_TYPE_F16 && + op->type == GGML_TYPE_F32 && + op->src[0]->ne[0] == 128 && + op->src[0]->ne[1] == 64 && + ggml_is_contiguous_rows(op->src[0]) && + ggml_is_contiguous_rows(op->src[1]) && + ggml_is_contiguous_rows(op->src[2]) && + ggml_is_contiguous_rows(op->src[3]); case GGML_OP_SSM_CONV: case GGML_OP_SSM_SCAN: return has_simdgroup_reduction; diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h index 330278d003d..d92950a5d02 100644 --- a/ggml/src/ggml-metal/ggml-metal-impl.h +++ b/ggml/src/ggml-metal/ggml-metal-impl.h @@ -1167,6 +1167,24 @@ typedef struct { int64_t val; } ggml_metal_kargs_memset; +typedef struct { + int32_t n_kv; + int32_t n_batch; + int32_t kv_offset; + int32_t mask_ne3; + uint64_t nb1; + uint64_t nb3; + uint64_t nbq1; + uint64_t nbq2; + uint64_t nbq3; + uint64_t nbk2; + uint64_t nbk3; + uint64_t nbw1; + uint64_t nbw3; + uint64_t nbm1; + uint64_t nbm3; +} ggml_metal_kargs_lightning_indexer; + typedef struct { int32_t ne00; int32_t ne01; diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index c716f118f6d..b8270f849da 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -316,6 +316,10 @@ static int ggml_metal_op_encode_impl(ggml_metal_op_t ctx, int idx) { { n_fuse = ggml_metal_op_cumsum(ctx, idx); } break; + case GGML_OP_LIGHTNING_INDEXER: + { + n_fuse = ggml_metal_op_lightning_indexer(ctx, idx); + } break; case GGML_OP_SOFT_MAX: { n_fuse = ggml_metal_op_soft_max(ctx, idx); @@ -1297,6 +1301,89 @@ int ggml_metal_op_diag(ggml_metal_op_t ctx, int idx) { return 1; } +int ggml_metal_op_lightning_indexer(ggml_metal_op_t ctx, int idx) { + ggml_tensor * op = ctx->node(idx); + + ggml_metal_encoder_t enc = ctx->enc; + + GGML_ASSERT(op->op == GGML_OP_LIGHTNING_INDEXER); + + const ggml_tensor * q = op->src[0]; + const ggml_tensor * k = op->src[1]; + const ggml_tensor * w = op->src[2]; + const ggml_tensor * m = op->src[3]; + + GGML_ASSERT(q->type == GGML_TYPE_F32); + GGML_ASSERT(k->type == GGML_TYPE_F16); + GGML_ASSERT(w->type == GGML_TYPE_F32); + GGML_ASSERT(m->type == GGML_TYPE_F16); + GGML_ASSERT(op->type == GGML_TYPE_F32); + + GGML_ASSERT(q->ne[0] == 128); + GGML_ASSERT(q->ne[1] == 64); + + ggml_metal_kargs_lightning_indexer args = { + /*.n_kv =*/ (int32_t) k->ne[2], + /*.n_batch =*/ (int32_t) q->ne[2], + /*.kv_offset =*/ 0, + /*.mask_ne3 =*/ (int32_t) m->ne[3], + /*.nb1 =*/ op->nb[1], + /*.nb3 =*/ op->nb[3], + /*.nbq1 =*/ q->nb[1], + /*.nbq2 =*/ q->nb[2], + /*.nbq3 =*/ q->nb[3], + /*.nbk2 =*/ k->nb[2], + /*.nbk3 =*/ k->nb[3], + /*.nbw1 =*/ w->nb[1], + /*.nbw3 =*/ w->nb[3], + /*.nbm1 =*/ m->nb[1], + /*.nbm3 =*/ m->nb[3], + }; + + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(q), 1); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(k), 2); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(w), 3); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(m), 4); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op), 5); + + constexpr int n_keys_simdgroup = 8; + constexpr int n_simdgroups = 8; + constexpr int n_keys_tg = n_keys_simdgroup*n_simdgroups; + constexpr int n_batch_tg = 8; + + const int32_t n_kv = args.n_kv; + const int32_t n_tg = n_kv/n_keys_tg; + int32_t kv_offset = 0; + + if (n_tg > 0) { + auto pipeline = ggml_metal_library_get_pipeline_lightning_indexer(ctx->lib, false); + ggml_metal_encoder_set_pipeline(enc, pipeline); + ggml_metal_encoder_set_bytes(enc, &args, sizeof(args), 0); + ggml_metal_encoder_dispatch_threadgroups(enc, n_tg, (q->ne[2] + n_batch_tg - 1)/n_batch_tg, q->ne[3], 32, n_simdgroups, 1); + kv_offset = n_tg*n_keys_tg; + } + + const int32_t n_simdgroups_tail = (n_kv - kv_offset)/n_keys_simdgroup; + if (n_simdgroups_tail > 0) { + args.kv_offset = kv_offset; + auto pipeline = ggml_metal_library_get_pipeline_lightning_indexer(ctx->lib, false); + ggml_metal_encoder_set_pipeline(enc, pipeline); + ggml_metal_encoder_set_bytes(enc, &args, sizeof(args), 0); + ggml_metal_encoder_dispatch_threadgroups(enc, 1, (q->ne[2] + n_batch_tg - 1)/n_batch_tg, q->ne[3], 32, n_simdgroups_tail, 1); + kv_offset += n_simdgroups_tail*n_keys_simdgroup; + } + + if (kv_offset < n_kv) { + args.kv_offset = kv_offset; + auto pipeline = ggml_metal_library_get_pipeline_lightning_indexer(ctx->lib, true); + ggml_metal_encoder_set_pipeline(enc, pipeline); + ggml_metal_encoder_set_bytes(enc, &args, sizeof(args), 0); + ggml_metal_encoder_dispatch_threadgroups(enc, 1, q->ne[2], q->ne[3], 32, 1, 1); + } + + return 1; +} + int ggml_metal_op_soft_max(ggml_metal_op_t ctx, int idx) { ggml_tensor * op = ctx->node(idx); diff --git a/ggml/src/ggml-metal/ggml-metal-ops.h b/ggml/src/ggml-metal/ggml-metal-ops.h index 89a6ad82f1c..79743147899 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.h +++ b/ggml/src/ggml-metal/ggml-metal-ops.h @@ -54,6 +54,7 @@ int ggml_metal_op_cumsum (ggml_metal_op_t ctx, int idx); int ggml_metal_op_get_rows (ggml_metal_op_t ctx, int idx); int ggml_metal_op_set_rows (ggml_metal_op_t ctx, int idx); int ggml_metal_op_diag (ggml_metal_op_t ctx, int idx); +int ggml_metal_op_lightning_indexer (ggml_metal_op_t ctx, int idx); int ggml_metal_op_soft_max (ggml_metal_op_t ctx, int idx); int ggml_metal_op_ssm_conv (ggml_metal_op_t ctx, int idx); int ggml_metal_op_ssm_scan (ggml_metal_op_t ctx, int idx); diff --git a/ggml/src/ggml-metal/ggml-metal.metal b/ggml/src/ggml-metal/ggml-metal.metal index 969fddfa5b8..256a475e13a 100644 --- a/ggml/src/ggml-metal/ggml-metal.metal +++ b/ggml/src/ggml-metal/ggml-metal.metal @@ -11216,3 +11216,144 @@ kernel void kernel_count_equal( typedef decltype(kernel_count_equal) kernel_count_equal_t; template [[host_name("kernel_count_equal_i32")]] kernel kernel_count_equal_t kernel_count_equal; + +constexpr constant short LI_N_EMBD = 128; +constexpr constant short LI_N_HEAD = 64; +constexpr constant short LI_N_KEY_SIMDGROUP = 8; +constexpr constant short LI_N_SIMDGROUP = 8; +constexpr constant short LI_N_BATCH_TG = 8; + +// Keep one 64-key K tile resident in simdgroup matrix registers while processing +// several queries. Q is staged eight heads at a time to preserve GPU occupancy. +kernel void kernel_lightning_indexer_f16( + constant ggml_metal_kargs_lightning_indexer & args, + device const char * q, + device const char * k, + device const char * w, + device const char * m, + device char * dst, + uint3 tgpig[[threadgroup_position_in_grid]], + ushort tiitg[[thread_index_in_threadgroup]], + ushort tiisg[[thread_index_in_simdgroup]], + ushort sgitg[[simdgroup_index_in_threadgroup]], + ushort3 ntg[[threads_per_threadgroup]]) { + constexpr short n_embd8 = LI_N_EMBD/8; + constexpr short n_embd4 = LI_N_EMBD/4; + constexpr short n_head_tile = 8; + constexpr short n_key_tg = LI_N_KEY_SIMDGROUP*LI_N_SIMDGROUP; + + const int i_stream = tgpig.z; + const int i_kv = args.kv_offset + tgpig.x*n_key_tg + sgitg*LI_N_KEY_SIMDGROUP; + + device const char * k_base = k + i_kv*args.nbk2 + i_stream*args.nbk3; + + simdgroup_half8x8 mk[n_embd8]; + + FOR_UNROLL (short i = 0; i < n_embd8; ++i) { + simdgroup_load(mk[i], (device const half *) k_base + 8*i, args.nbk2/sizeof(half), 0, true); + } + + threadgroup half q_shared[n_head_tile*LI_N_EMBD]; + threadgroup float w_shared[n_head_tile]; + threadgroup float qk_shared[LI_N_SIMDGROUP*n_head_tile*LI_N_KEY_SIMDGROUP]; + + const int i_batch_0 = tgpig.y*LI_N_BATCH_TG; + const int n_batch = min((int) LI_N_BATCH_TG, args.n_batch - i_batch_0); + + for (short ib = 0; ib < n_batch; ++ib) { + const int i_batch = i_batch_0 + ib; + device const char * q_base = q + i_batch*args.nbq2 + i_stream*args.nbq3; + device const char * w_base = w + i_batch*args.nbw1 + i_stream*args.nbw3; + + float score = 0.0f; + + FOR_UNROLL (short i_head = 0; i_head < LI_N_HEAD; i_head += n_head_tile) { + for (short i = tiitg; i < n_head_tile*n_embd4; i += ntg.x*ntg.y) { + const short ih = i/n_embd4; + const short i4 = i - ih*n_embd4; + device const float4 * q4 = (device const float4 *) (q_base + (i_head + ih)*args.nbq1); + *(threadgroup half4 *) (q_shared + ih*LI_N_EMBD + 4*i4) = half4(q4[i4]); + } + + if (tiitg < n_head_tile) { + w_shared[tiitg] = *((device const float *) (w_base + (i_head + tiitg)*sizeof(float))); + } + + threadgroup_barrier(mem_flags::mem_threadgroup); + + simdgroup_float8x8 mqk = make_filled_simdgroup_matrix(0.0f); + + FOR_UNROLL (short i = 0; i < n_embd8; ++i) { + simdgroup_half8x8 mq; + simdgroup_load(mq, q_shared + 8*i, LI_N_EMBD, 0, false); + simdgroup_multiply_accumulate(mqk, mq, mk[i], mqk); + } + + threadgroup float * qk = qk_shared + sgitg*n_head_tile*LI_N_KEY_SIMDGROUP; + simdgroup_store(mqk, qk, LI_N_KEY_SIMDGROUP, 0, false); + simdgroup_barrier(mem_flags::mem_threadgroup); + + if (tiisg < LI_N_KEY_SIMDGROUP) { + FOR_UNROLL (short ih = 0; ih < n_head_tile; ++ih) { + score += max(qk[ih*LI_N_KEY_SIMDGROUP + tiisg], 0.0f)*w_shared[ih]; + } + } + + threadgroup_barrier(mem_flags::mem_threadgroup); + } + + if (tiisg < LI_N_KEY_SIMDGROUP) { + const int ik = i_kv + tiisg; + device const half * mask = (device const half *) (m + i_batch*args.nbm1 + (i_stream % args.mask_ne3)*args.nbm3); + device float * out = (device float *) (dst + i_batch*args.nb1 + i_stream*args.nb3); + out[ik] = score + float(mask[ik]); + } + } +} + +kernel void kernel_lightning_indexer_f16_tail( + constant ggml_metal_kargs_lightning_indexer & args, + device const char * q, + device const char * k, + device const char * w, + device const char * m, + device char * dst, + uint3 tgpig[[threadgroup_position_in_grid]], + ushort tiisg[[thread_index_in_simdgroup]]) { + const int i_batch = tgpig.y; + const int i_stream = tgpig.z; + const int n_key = args.n_kv - args.kv_offset; + + device const char * q_base = q + i_batch*args.nbq2 + i_stream*args.nbq3; + device const char * k_base = k + args.kv_offset*args.nbk2 + i_stream*args.nbk3; + device const char * w_base = w + i_batch*args.nbw1 + i_stream*args.nbw3; + + float4 k4[LI_N_KEY_SIMDGROUP]; + float score[LI_N_KEY_SIMDGROUP] = { 0.0f }; + + for (short ik = 0; ik < n_key; ++ik) { + device const half4 * row = (device const half4 *) (k_base + ik*args.nbk2); + k4[ik] = float4(row[tiisg]); + } + + FOR_UNROLL (short ih = 0; ih < LI_N_HEAD; ++ih) { + device const float4 * row = (device const float4 *) (q_base + ih*args.nbq1); + const float4 q4 = row[tiisg]; + const float weight = *((device const float *) (w_base + ih*sizeof(float))); + + for (short ik = 0; ik < n_key; ++ik) { + float qk = simd_sum(dot(q4, k4[ik])); + score[ik] += max(qk, 0.0f)*weight; + } + } + + if (tiisg == 0) { + device const half * mask = (device const half *) (m + i_batch*args.nbm1 + (i_stream % args.mask_ne3)*args.nbm3); + device float * out = (device float *) (dst + i_batch*args.nb1 + i_stream*args.nb3); + + for (short ik = 0; ik < n_key; ++ik) { + const int i = args.kv_offset + ik; + out[i] = score[ik] + float(mask[i]); + } + } +} diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index e3553c5e9ca..2fe5faf0b6e 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -9668,6 +9668,10 @@ static std::vector> make_test_cases_eval() { } } + for (int kv : { 1, 7, 8, 63, 64, 65 }) { + test_cases.emplace_back(new test_lightning_indexer(128, 64, kv, 1, 4, 1, GGML_TYPE_F16)); + } + return test_cases; } #ifdef _MSC_VER From 54e6fdc2a09af667c7d888cf3bbd24f4b64ee06c Mon Sep 17 00:00:00 2001 From: Thiago Padilha Date: Sun, 2 Aug 2026 06:59:59 -0300 Subject: [PATCH 2/4] metal: stage Lightning Indexer K tiles - Stage and dequantize K in F16 threadgroup memory before simdgroup matrix loads. - Zero-fill partial tiles and guard stores so all KV segments use the same numerical path. - Support F32, F16, BF16, Q4_0, Q4_1, Q5_0, Q5_1, and Q8_0 K caches. llama-bench (--mmap 1, -fa on, -p 512, -n 128; d=0/10k/20k): - pp512: 160.38 +/- 1.01 t/s - tg128: 9.08 +/- 0.03 t/s - pp512 @ d10000: 88.37 +/- 0.46 t/s - tg128 @ d10000: 9.07 +/- 0.04 t/s - pp512 @ d20000: 62.53 +/- 0.46 t/s - tg128 @ d20000: 8.84 +/- 0.03 t/s Assisted-by: Codex --- ggml/src/ggml-metal/ggml-metal-device.cpp | 10 ++- ggml/src/ggml-metal/ggml-metal-device.h | 2 +- ggml/src/ggml-metal/ggml-metal-device.m | 39 ++++++--- ggml/src/ggml-metal/ggml-metal-impl.h | 1 - ggml/src/ggml-metal/ggml-metal-ops.cpp | 53 +++++------- ggml/src/ggml-metal/ggml-metal.metal | 99 +++++++++++------------ tests/test-backend-ops.cpp | 4 +- 7 files changed, 104 insertions(+), 104 deletions(-) diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp index f265f0739d8..ea53306beea 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.cpp +++ b/ggml/src/ggml-metal/ggml-metal-device.cpp @@ -477,8 +477,14 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_soft_max(ggml_me return res; } -ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning_indexer(ggml_metal_library_t lib, bool tail) { - const char * name = tail ? "kernel_lightning_indexer_f16_tail" : "kernel_lightning_indexer_f16"; +ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning_indexer( + ggml_metal_library_t lib, + const ggml_tensor * op) { + GGML_ASSERT(op->op == GGML_OP_LIGHTNING_INDEXER); + + char name[256]; + + snprintf(name, 256, "kernel_lightning_indexer_%s", ggml_type_name(op->src[1]->type)); ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name); if (!res.pipeline) { diff --git a/ggml/src/ggml-metal/ggml-metal-device.h b/ggml/src/ggml-metal/ggml-metal-device.h index f4262ddd0a7..7df110cbd0e 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.h +++ b/ggml/src/ggml-metal/ggml-metal-device.h @@ -124,7 +124,7 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_cumsum_bl struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_cumsum_add (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_tri (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_soft_max (ggml_metal_library_t lib, const struct ggml_tensor * op); -struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning_indexer (ggml_metal_library_t lib, bool tail); +struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning_indexer (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv_batched (ggml_metal_library_t lib, const struct ggml_tensor * op, int ssm_conv_bs); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan (ggml_metal_library_t lib, const struct ggml_tensor * op); diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index de69b694785..bbd121fbc39 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -1300,18 +1300,33 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te } return has_simdgroup_mm; // TODO: over-restricted for vec-kernels case GGML_OP_LIGHTNING_INDEXER: - return has_simdgroup_mm && - op->src[0]->type == GGML_TYPE_F32 && - op->src[1]->type == GGML_TYPE_F16 && - op->src[2]->type == GGML_TYPE_F32 && - op->src[3]->type == GGML_TYPE_F16 && - op->type == GGML_TYPE_F32 && - op->src[0]->ne[0] == 128 && - op->src[0]->ne[1] == 64 && - ggml_is_contiguous_rows(op->src[0]) && - ggml_is_contiguous_rows(op->src[1]) && - ggml_is_contiguous_rows(op->src[2]) && - ggml_is_contiguous_rows(op->src[3]); + if (!has_simdgroup_mm || + op->src[0]->type != GGML_TYPE_F32 || + op->src[2]->type != GGML_TYPE_F32 || + op->src[3]->type != GGML_TYPE_F16 || + op->type != GGML_TYPE_F32 || + op->src[0]->ne[0] != 128 || + op->src[0]->ne[1] != 64 || + !ggml_is_contiguous_rows(op->src[0]) || + !ggml_is_contiguous_rows(op->src[1]) || + !ggml_is_contiguous_rows(op->src[2]) || + !ggml_is_contiguous_rows(op->src[3])) { + return false; + } + switch (op->src[1]->type) { + case GGML_TYPE_F32: + case GGML_TYPE_F16: + case GGML_TYPE_Q4_0: + case GGML_TYPE_Q4_1: + case GGML_TYPE_Q5_0: + case GGML_TYPE_Q5_1: + case GGML_TYPE_Q8_0: + return true; + case GGML_TYPE_BF16: + return has_bfloat; + default: + return false; + } case GGML_OP_SSM_CONV: case GGML_OP_SSM_SCAN: return has_simdgroup_reduction; diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h index d7d6f237ae2..18227349f1f 100644 --- a/ggml/src/ggml-metal/ggml-metal-impl.h +++ b/ggml/src/ggml-metal/ggml-metal-impl.h @@ -1174,7 +1174,6 @@ typedef struct { typedef struct { int32_t n_kv; int32_t n_batch; - int32_t kv_offset; int32_t mask_ne3; uint64_t nb1; uint64_t nb3; diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index b8986c09ffb..274a3eb9e45 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -1314,7 +1314,14 @@ int ggml_metal_op_lightning_indexer(ggml_metal_op_t ctx, int idx) { const ggml_tensor * m = op->src[3]; GGML_ASSERT(q->type == GGML_TYPE_F32); - GGML_ASSERT(k->type == GGML_TYPE_F16); + GGML_ASSERT(k->type == GGML_TYPE_F32 || + k->type == GGML_TYPE_F16 || + k->type == GGML_TYPE_BF16 || + k->type == GGML_TYPE_Q4_0 || + k->type == GGML_TYPE_Q4_1 || + k->type == GGML_TYPE_Q5_0 || + k->type == GGML_TYPE_Q5_1 || + k->type == GGML_TYPE_Q8_0); GGML_ASSERT(w->type == GGML_TYPE_F32); GGML_ASSERT(m->type == GGML_TYPE_F16); GGML_ASSERT(op->type == GGML_TYPE_F32); @@ -1325,7 +1332,6 @@ int ggml_metal_op_lightning_indexer(ggml_metal_op_t ctx, int idx) { ggml_metal_kargs_lightning_indexer args = { /*.n_kv =*/ (int32_t) k->ne[2], /*.n_batch =*/ (int32_t) q->ne[2], - /*.kv_offset =*/ 0, /*.mask_ne3 =*/ (int32_t) m->ne[3], /*.nb1 =*/ op->nb[1], /*.nb3 =*/ op->nb[3], @@ -1346,40 +1352,17 @@ int ggml_metal_op_lightning_indexer(ggml_metal_op_t ctx, int idx) { ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(m), 4); ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op), 5); - constexpr int n_keys_simdgroup = 8; - constexpr int n_simdgroups = 8; - constexpr int n_keys_tg = n_keys_simdgroup*n_simdgroups; - constexpr int n_batch_tg = 8; + constexpr int n_keys_tg = 64; + constexpr int n_batch_tg = 8; + constexpr int nsg = 8; - const int32_t n_kv = args.n_kv; - const int32_t n_tg = n_kv/n_keys_tg; - int32_t kv_offset = 0; - - if (n_tg > 0) { - auto pipeline = ggml_metal_library_get_pipeline_lightning_indexer(ctx->lib, false); - ggml_metal_encoder_set_pipeline(enc, pipeline); - ggml_metal_encoder_set_bytes(enc, &args, sizeof(args), 0); - ggml_metal_encoder_dispatch_threadgroups(enc, n_tg, (q->ne[2] + n_batch_tg - 1)/n_batch_tg, q->ne[3], 32, n_simdgroups, 1); - kv_offset = n_tg*n_keys_tg; - } - - const int32_t n_simdgroups_tail = (n_kv - kv_offset)/n_keys_simdgroup; - if (n_simdgroups_tail > 0) { - args.kv_offset = kv_offset; - auto pipeline = ggml_metal_library_get_pipeline_lightning_indexer(ctx->lib, false); - ggml_metal_encoder_set_pipeline(enc, pipeline); - ggml_metal_encoder_set_bytes(enc, &args, sizeof(args), 0); - ggml_metal_encoder_dispatch_threadgroups(enc, 1, (q->ne[2] + n_batch_tg - 1)/n_batch_tg, q->ne[3], 32, n_simdgroups_tail, 1); - kv_offset += n_simdgroups_tail*n_keys_simdgroup; - } - - if (kv_offset < n_kv) { - args.kv_offset = kv_offset; - auto pipeline = ggml_metal_library_get_pipeline_lightning_indexer(ctx->lib, true); - ggml_metal_encoder_set_pipeline(enc, pipeline); - ggml_metal_encoder_set_bytes(enc, &args, sizeof(args), 0); - ggml_metal_encoder_dispatch_threadgroups(enc, 1, q->ne[2], q->ne[3], 32, 1, 1); - } + auto pipeline = ggml_metal_library_get_pipeline_lightning_indexer(ctx->lib, op); + ggml_metal_encoder_set_pipeline(enc, pipeline); + ggml_metal_encoder_set_bytes(enc, &args, sizeof(args), 0); + ggml_metal_encoder_dispatch_threadgroups(enc, + (k->ne[2] + n_keys_tg - 1)/n_keys_tg, + (q->ne[2] + n_batch_tg - 1)/n_batch_tg, + q->ne[3], 32, nsg, 1); return 1; } diff --git a/ggml/src/ggml-metal/ggml-metal.metal b/ggml/src/ggml-metal/ggml-metal.metal index 63e4e650dfb..78b769d7ec2 100644 --- a/ggml/src/ggml-metal/ggml-metal.metal +++ b/ggml/src/ggml-metal/ggml-metal.metal @@ -11285,9 +11285,11 @@ constexpr constant short LI_N_KEY_SIMDGROUP = 8; constexpr constant short LI_N_SIMDGROUP = 8; constexpr constant short LI_N_BATCH_TG = 8; -// Keep one 64-key K tile resident in simdgroup matrix registers while processing -// several queries. Q is staged eight heads at a time to preserve GPU occupancy. -kernel void kernel_lightning_indexer_f16( +template< + typename kd4x4_t, + short nl_k, + void (*deq_k)(device const kd4x4_t *, short, thread half4x4 &)> +kernel void kernel_lightning_indexer( constant ggml_metal_kargs_lightning_indexer & args, device const char * q, device const char * k, @@ -11300,19 +11302,42 @@ kernel void kernel_lightning_indexer_f16( ushort sgitg[[simdgroup_index_in_threadgroup]], ushort3 ntg[[threads_per_threadgroup]]) { constexpr short n_embd8 = LI_N_EMBD/8; + constexpr short n_embd16 = LI_N_EMBD/16; constexpr short n_embd4 = LI_N_EMBD/4; constexpr short n_head_tile = 8; constexpr short n_key_tg = LI_N_KEY_SIMDGROUP*LI_N_SIMDGROUP; const int i_stream = tgpig.z; - const int i_kv = args.kv_offset + tgpig.x*n_key_tg + sgitg*LI_N_KEY_SIMDGROUP; + const int i_kv_0 = tgpig.x*n_key_tg; + const int i_kv = i_kv_0 + sgitg*LI_N_KEY_SIMDGROUP; + + threadgroup half k_shared[n_key_tg*LI_N_EMBD]; + threadgroup half4x4 * k4x4_shared = (threadgroup half4x4 *) k_shared; - device const char * k_base = k + i_kv*args.nbk2 + i_stream*args.nbk3; + for (short i = tiitg; i < n_key_tg*n_embd16; i += ntg.x*ntg.y) { + const short ik = i/n_embd16; + const short i16 = i - ik*n_embd16; + const int i_k = i_kv_0 + ik; + + half4x4 tmp; + if (i_k < args.n_kv) { + device const kd4x4_t * row = (device const kd4x4_t *) (k + i_k*args.nbk2 + i_stream*args.nbk3); + deq_k(row + i16/nl_k, i16%nl_k, tmp); + } else { + FOR_UNROLL (short j = 0; j < 4; ++j) { + tmp[j] = half4(0.0h); + } + } + + k4x4_shared[i] = tmp; + } + + threadgroup_barrier(mem_flags::mem_threadgroup); simdgroup_half8x8 mk[n_embd8]; FOR_UNROLL (short i = 0; i < n_embd8; ++i) { - simdgroup_load(mk[i], (device const half *) k_base + 8*i, args.nbk2/sizeof(half), 0, true); + simdgroup_load(mk[i], k_shared + sgitg*LI_N_KEY_SIMDGROUP*LI_N_EMBD + 8*i, LI_N_EMBD, 0, true); } threadgroup half q_shared[n_head_tile*LI_N_EMBD]; @@ -11366,56 +11391,26 @@ kernel void kernel_lightning_indexer_f16( if (tiisg < LI_N_KEY_SIMDGROUP) { const int ik = i_kv + tiisg; - device const half * mask = (device const half *) (m + i_batch*args.nbm1 + (i_stream % args.mask_ne3)*args.nbm3); - device float * out = (device float *) (dst + i_batch*args.nb1 + i_stream*args.nb3); - out[ik] = score + float(mask[ik]); + if (ik < args.n_kv) { + device const half * mask = (device const half *) (m + i_batch*args.nbm1 + (i_stream % args.mask_ne3)*args.nbm3); + device float * out = (device float *) (dst + i_batch*args.nb1 + i_stream*args.nb3); + out[ik] = score + float(mask[ik]); + } } } } -kernel void kernel_lightning_indexer_f16_tail( - constant ggml_metal_kargs_lightning_indexer & args, - device const char * q, - device const char * k, - device const char * w, - device const char * m, - device char * dst, - uint3 tgpig[[threadgroup_position_in_grid]], - ushort tiisg[[thread_index_in_simdgroup]]) { - const int i_batch = tgpig.y; - const int i_stream = tgpig.z; - const int n_key = args.n_kv - args.kv_offset; - - device const char * q_base = q + i_batch*args.nbq2 + i_stream*args.nbq3; - device const char * k_base = k + args.kv_offset*args.nbk2 + i_stream*args.nbk3; - device const char * w_base = w + i_batch*args.nbw1 + i_stream*args.nbw3; +typedef decltype(kernel_lightning_indexer) kernel_lightning_indexer_t; - float4 k4[LI_N_KEY_SIMDGROUP]; - float score[LI_N_KEY_SIMDGROUP] = { 0.0f }; +template [[host_name("kernel_lightning_indexer_f32")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer; +template [[host_name("kernel_lightning_indexer_f16")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer; - for (short ik = 0; ik < n_key; ++ik) { - device const half4 * row = (device const half4 *) (k_base + ik*args.nbk2); - k4[ik] = float4(row[tiisg]); - } - - FOR_UNROLL (short ih = 0; ih < LI_N_HEAD; ++ih) { - device const float4 * row = (device const float4 *) (q_base + ih*args.nbq1); - const float4 q4 = row[tiisg]; - const float weight = *((device const float *) (w_base + ih*sizeof(float))); - - for (short ik = 0; ik < n_key; ++ik) { - float qk = simd_sum(dot(q4, k4[ik])); - score[ik] += max(qk, 0.0f)*weight; - } - } - - if (tiisg == 0) { - device const half * mask = (device const half *) (m + i_batch*args.nbm1 + (i_stream % args.mask_ne3)*args.nbm3); - device float * out = (device float *) (dst + i_batch*args.nb1 + i_stream*args.nb3); +#if defined(GGML_METAL_HAS_BF16) +template [[host_name("kernel_lightning_indexer_bf16")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer; +#endif - for (short ik = 0; ik < n_key; ++ik) { - const int i = args.kv_offset + ik; - out[i] = score[ik] + float(mask[i]); - } - } -} +template [[host_name("kernel_lightning_indexer_q4_0")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer; +template [[host_name("kernel_lightning_indexer_q4_1")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer; +template [[host_name("kernel_lightning_indexer_q5_0")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer; +template [[host_name("kernel_lightning_indexer_q5_1")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer; +template [[host_name("kernel_lightning_indexer_q8_0")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer; diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 668065f5992..575edf9331e 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -9729,7 +9729,9 @@ static std::vector> make_test_cases_eval() { } for (int kv : { 1, 7, 8, 63, 64, 65 }) { - test_cases.emplace_back(new test_lightning_indexer(128, 64, kv, 1, 4, 1, GGML_TYPE_F16)); + for (ggml_type type_K : {GGML_TYPE_F32, GGML_TYPE_F16, GGML_TYPE_BF16, GGML_TYPE_Q8_0, GGML_TYPE_Q5_1, GGML_TYPE_Q5_0, GGML_TYPE_Q4_1, GGML_TYPE_Q4_0}) { + test_cases.emplace_back(new test_lightning_indexer(128, 64, kv, 1, 4, 1, type_K)); + } } return test_cases; From 280158e1740dbcc8b6f7ea42c7dd02ea4b9048a8 Mon Sep 17 00:00:00 2001 From: forforever73 <690105611@qq.com> Date: Mon, 3 Aug 2026 00:02:21 +0800 Subject: [PATCH 3/4] dedup Lightning Indexer constants, fix flaky test --- ggml/src/ggml-metal/ggml-metal-device.m | 7 +- ggml/src/ggml-metal/ggml-metal-impl.h | 7 ++ ggml/src/ggml-metal/ggml-metal-ops.cpp | 14 +-- ggml/src/ggml-metal/ggml-metal.metal | 127 +++++++++++++----------- tests/test-backend-ops.cpp | 2 +- 5 files changed, 90 insertions(+), 67 deletions(-) diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index bbd121fbc39..69f63a48f47 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -2,6 +2,7 @@ #import "ggml-impl.h" #import "ggml-backend-impl.h" +#import "ggml-metal-impl.h" #include @@ -1300,13 +1301,15 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te } return has_simdgroup_mm; // TODO: over-restricted for vec-kernels case GGML_OP_LIGHTNING_INDEXER: + if (op->src[0]->ne[0] != OP_LIGHTNING_INDEXER_DK || + op->src[0]->ne[1] != OP_LIGHTNING_INDEXER_NH) { + return false; + } if (!has_simdgroup_mm || op->src[0]->type != GGML_TYPE_F32 || op->src[2]->type != GGML_TYPE_F32 || op->src[3]->type != GGML_TYPE_F16 || op->type != GGML_TYPE_F32 || - op->src[0]->ne[0] != 128 || - op->src[0]->ne[1] != 64 || !ggml_is_contiguous_rows(op->src[0]) || !ggml_is_contiguous_rows(op->src[1]) || !ggml_is_contiguous_rows(op->src[2]) || diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h index 18227349f1f..5d9d016d4ba 100644 --- a/ggml/src/ggml-metal/ggml-metal-impl.h +++ b/ggml/src/ggml-metal/ggml-metal-impl.h @@ -112,6 +112,13 @@ #define OP_FLASH_ATTN_EXT_VEC_NQPSG 1 #define OP_FLASH_ATTN_EXT_VEC_NCPSG 32 +#define OP_LIGHTNING_INDEXER_DK 128 +#define OP_LIGHTNING_INDEXER_NH 64 +#define OP_LIGHTNING_INDEXER_NHPTG 8 +#define OP_LIGHTNING_INDEXER_NKPSG 8 +#define OP_LIGHTNING_INDEXER_NSG 8 +#define OP_LIGHTNING_INDEXER_NBPTG 8 + #define OP_UNARY_NUM_SCALE 10 #define OP_UNARY_NUM_FILL 11 #define OP_UNARY_NUM_CLAMP 12 diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index 274a3eb9e45..5109bb0f977 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -1326,8 +1326,8 @@ int ggml_metal_op_lightning_indexer(ggml_metal_op_t ctx, int idx) { GGML_ASSERT(m->type == GGML_TYPE_F16); GGML_ASSERT(op->type == GGML_TYPE_F32); - GGML_ASSERT(q->ne[0] == 128); - GGML_ASSERT(q->ne[1] == 64); + GGML_ASSERT(q->ne[0] == OP_LIGHTNING_INDEXER_DK); + GGML_ASSERT(q->ne[1] == OP_LIGHTNING_INDEXER_NH); ggml_metal_kargs_lightning_indexer args = { /*.n_kv =*/ (int32_t) k->ne[2], @@ -1352,16 +1352,16 @@ int ggml_metal_op_lightning_indexer(ggml_metal_op_t ctx, int idx) { ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(m), 4); ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op), 5); - constexpr int n_keys_tg = 64; - constexpr int n_batch_tg = 8; - constexpr int nsg = 8; + const int nsg = OP_LIGHTNING_INDEXER_NSG; + const int nkptg = OP_LIGHTNING_INDEXER_NKPSG*nsg; + const int nbptg = OP_LIGHTNING_INDEXER_NBPTG; auto pipeline = ggml_metal_library_get_pipeline_lightning_indexer(ctx->lib, op); ggml_metal_encoder_set_pipeline(enc, pipeline); ggml_metal_encoder_set_bytes(enc, &args, sizeof(args), 0); ggml_metal_encoder_dispatch_threadgroups(enc, - (k->ne[2] + n_keys_tg - 1)/n_keys_tg, - (q->ne[2] + n_batch_tg - 1)/n_batch_tg, + (k->ne[2] + nkptg - 1)/nkptg, + (q->ne[2] + nbptg - 1)/nbptg, q->ne[3], 32, nsg, 1); return 1; diff --git a/ggml/src/ggml-metal/ggml-metal.metal b/ggml/src/ggml-metal/ggml-metal.metal index 78b769d7ec2..8d4e42c314e 100644 --- a/ggml/src/ggml-metal/ggml-metal.metal +++ b/ggml/src/ggml-metal/ggml-metal.metal @@ -11279,12 +11279,6 @@ typedef decltype(kernel_count_equal) kernel_count_equal_t; template [[host_name("kernel_count_equal_i32")]] kernel kernel_count_equal_t kernel_count_equal; -constexpr constant short LI_N_EMBD = 128; -constexpr constant short LI_N_HEAD = 64; -constexpr constant short LI_N_KEY_SIMDGROUP = 8; -constexpr constant short LI_N_SIMDGROUP = 8; -constexpr constant short LI_N_BATCH_TG = 8; - template< typename kd4x4_t, short nl_k, @@ -11296,105 +11290,124 @@ kernel void kernel_lightning_indexer( device const char * w, device const char * m, device char * dst, - uint3 tgpig[[threadgroup_position_in_grid]], - ushort tiitg[[thread_index_in_threadgroup]], - ushort tiisg[[thread_index_in_simdgroup]], - ushort sgitg[[simdgroup_index_in_threadgroup]], - ushort3 ntg[[threads_per_threadgroup]]) { - constexpr short n_embd8 = LI_N_EMBD/8; - constexpr short n_embd16 = LI_N_EMBD/16; - constexpr short n_embd4 = LI_N_EMBD/4; - constexpr short n_head_tile = 8; - constexpr short n_key_tg = LI_N_KEY_SIMDGROUP*LI_N_SIMDGROUP; + uint3 tgpig[[threadgroup_position_in_grid]], + ushort tiitg[[thread_index_in_threadgroup]], + ushort tiisg[[thread_index_in_simdgroup]], + ushort sgitg[[simdgroup_index_in_threadgroup]]) { + constexpr short DK = OP_LIGHTNING_INDEXER_DK; + constexpr short NH = OP_LIGHTNING_INDEXER_NH; + constexpr short NHPTG = OP_LIGHTNING_INDEXER_NHPTG; + constexpr short NKPSG = OP_LIGHTNING_INDEXER_NKPSG; + constexpr short NSG = OP_LIGHTNING_INDEXER_NSG; + constexpr short NBPTG = OP_LIGHTNING_INDEXER_NBPTG; + + constexpr short DK4 = DK/4; + constexpr short DK8 = DK/8; + constexpr short DK16 = DK/16; + + constexpr short NK = NKPSG*NSG; // keys per threadgroup + constexpr short NTG = 32*NSG; // threads per threadgroup const int i_stream = tgpig.z; - const int i_kv_0 = tgpig.x*n_key_tg; - const int i_kv = i_kv_0 + sgitg*LI_N_KEY_SIMDGROUP; + const int i_kv_0 = tgpig.x*NK; // first key of this threadgroup + const int i_kv = i_kv_0 + sgitg*NKPSG; // first key of this simdgroup - threadgroup half k_shared[n_key_tg*LI_N_EMBD]; - threadgroup half4x4 * k4x4_shared = (threadgroup half4x4 *) k_shared; + threadgroup half4x4 sk4x4[NK*DK16]; + threadgroup half * sk = (threadgroup half *) sk4x4; - for (short i = tiitg; i < n_key_tg*n_embd16; i += ntg.x*ntg.y) { - const short ik = i/n_embd16; - const short i16 = i - ik*n_embd16; - const int i_k = i_kv_0 + ik; + for (short i = tiitg; i < NK*DK16; i += NTG) { + const short ik = i/DK16; + const short i16 = i%DK16; half4x4 tmp; - if (i_k < args.n_kv) { - device const kd4x4_t * row = (device const kd4x4_t *) (k + i_k*args.nbk2 + i_stream*args.nbk3); - deq_k(row + i16/nl_k, i16%nl_k, tmp); + + if (i_kv_0 + ik < args.n_kv) { + device const kd4x4_t * kr = (device const kd4x4_t *) (k + (i_kv_0 + ik)*args.nbk2 + i_stream*args.nbk3); + + deq_k(kr + i16/nl_k, i16%nl_k, tmp); } else { FOR_UNROLL (short j = 0; j < 4; ++j) { tmp[j] = half4(0.0h); } } - k4x4_shared[i] = tmp; + sk4x4[i] = tmp; } threadgroup_barrier(mem_flags::mem_threadgroup); - simdgroup_half8x8 mk[n_embd8]; + // K tile of this simdgroup, transposed to [DK, NKPSG] + simdgroup_half8x8 mk[DK8]; - FOR_UNROLL (short i = 0; i < n_embd8; ++i) { - simdgroup_load(mk[i], k_shared + sgitg*LI_N_KEY_SIMDGROUP*LI_N_EMBD + 8*i, LI_N_EMBD, 0, true); + FOR_UNROLL (short i = 0; i < DK8; ++i) { + simdgroup_load(mk[i], sk + sgitg*NKPSG*DK + 8*i, DK, 0, true); } - threadgroup half q_shared[n_head_tile*LI_N_EMBD]; - threadgroup float w_shared[n_head_tile]; - threadgroup float qk_shared[LI_N_SIMDGROUP*n_head_tile*LI_N_KEY_SIMDGROUP]; + threadgroup half4 sq4[NHPTG*DK4]; + threadgroup half * sq = (threadgroup half *) sq4; + + threadgroup float sw [NHPTG]; + threadgroup float sqk[NSG*NHPTG*NKPSG]; - const int i_batch_0 = tgpig.y*LI_N_BATCH_TG; - const int n_batch = min((int) LI_N_BATCH_TG, args.n_batch - i_batch_0); + const int i_batch_0 = tgpig.y*NBPTG; + const int n_batch = min((int) NBPTG, args.n_batch - i_batch_0); for (short ib = 0; ib < n_batch; ++ib) { const int i_batch = i_batch_0 + ib; - device const char * q_base = q + i_batch*args.nbq2 + i_stream*args.nbq3; - device const char * w_base = w + i_batch*args.nbw1 + i_stream*args.nbw3; + + device const char * pq = q + i_batch*args.nbq2 + i_stream*args.nbq3; + device const char * pw = w + i_batch*args.nbw1 + i_stream*args.nbw3; float score = 0.0f; - FOR_UNROLL (short i_head = 0; i_head < LI_N_HEAD; i_head += n_head_tile) { - for (short i = tiitg; i < n_head_tile*n_embd4; i += ntg.x*ntg.y) { - const short ih = i/n_embd4; - const short i4 = i - ih*n_embd4; - device const float4 * q4 = (device const float4 *) (q_base + (i_head + ih)*args.nbq1); - *(threadgroup half4 *) (q_shared + ih*LI_N_EMBD + 4*i4) = half4(q4[i4]); + FOR_UNROLL (short i_head = 0; i_head < NH; i_head += NHPTG) { + // stage the Q tile [DK, NHPTG] and the (prescaled) head weights + for (short i = tiitg; i < NHPTG*DK4; i += NTG) { + const short ih = i/DK4; + const short i4 = i%DK4; + + device const float4 * q4 = (device const float4 *) (pq + (i_head + ih)*args.nbq1); + + sq4[ih*DK4 + i4] = half4(q4[i4]); } - if (tiitg < n_head_tile) { - w_shared[tiitg] = *((device const float *) (w_base + (i_head + tiitg)*sizeof(float))); + if (tiitg < NHPTG) { + sw[tiitg] = ((device const float *) pw)[i_head + tiitg]; } threadgroup_barrier(mem_flags::mem_threadgroup); simdgroup_float8x8 mqk = make_filled_simdgroup_matrix(0.0f); - FOR_UNROLL (short i = 0; i < n_embd8; ++i) { + FOR_UNROLL (short i = 0; i < DK8; ++i) { simdgroup_half8x8 mq; - simdgroup_load(mq, q_shared + 8*i, LI_N_EMBD, 0, false); + + simdgroup_load(mq, sq + 8*i, DK, 0, false); simdgroup_multiply_accumulate(mqk, mq, mk[i], mqk); } - threadgroup float * qk = qk_shared + sgitg*n_head_tile*LI_N_KEY_SIMDGROUP; - simdgroup_store(mqk, qk, LI_N_KEY_SIMDGROUP, 0, false); + threadgroup float * pqk = sqk + sgitg*NHPTG*NKPSG; + + simdgroup_store(mqk, pqk, NKPSG, 0, false); simdgroup_barrier(mem_flags::mem_threadgroup); - if (tiisg < LI_N_KEY_SIMDGROUP) { - FOR_UNROLL (short ih = 0; ih < n_head_tile; ++ih) { - score += max(qk[ih*LI_N_KEY_SIMDGROUP + tiisg], 0.0f)*w_shared[ih]; + // one lane per key: ReLU, apply the head weight and accumulate over the head tile + if (tiisg < NKPSG) { + FOR_UNROLL (short ih = 0; ih < NHPTG; ++ih) { + score += max(pqk[ih*NKPSG + tiisg], 0.0f)*sw[ih]; } } threadgroup_barrier(mem_flags::mem_threadgroup); } - if (tiisg < LI_N_KEY_SIMDGROUP) { + if (tiisg < NKPSG) { const int ik = i_kv + tiisg; if (ik < args.n_kv) { - device const half * mask = (device const half *) (m + i_batch*args.nbm1 + (i_stream % args.mask_ne3)*args.nbm3); - device float * out = (device float *) (dst + i_batch*args.nb1 + i_stream*args.nb3); - out[ik] = score + float(mask[ik]); + device const half * pm = (device const half *) (m + i_batch*args.nbm1 + (i_stream % args.mask_ne3)*args.nbm3); + device float * pd = (device float *) (dst + i_batch*args.nb1 + i_stream*args.nb3); + + pd[ik] = score + (float) pm[ik]; } } } diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 575edf9331e..ccd09ee29c1 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -9730,7 +9730,7 @@ static std::vector> make_test_cases_eval() { for (int kv : { 1, 7, 8, 63, 64, 65 }) { for (ggml_type type_K : {GGML_TYPE_F32, GGML_TYPE_F16, GGML_TYPE_BF16, GGML_TYPE_Q8_0, GGML_TYPE_Q5_1, GGML_TYPE_Q5_0, GGML_TYPE_Q4_1, GGML_TYPE_Q4_0}) { - test_cases.emplace_back(new test_lightning_indexer(128, 64, kv, 1, 4, 1, type_K)); + test_cases.emplace_back(new test_lightning_indexer(128, 64, kv, 32, 4, 1, type_K)); } } From be99a4fc79356ab8134b9a818fe8e59528688b6c Mon Sep 17 00:00:00 2001 From: Georgi Gerganov Date: Mon, 3 Aug 2026 07:30:56 +0300 Subject: [PATCH 4/4] cont : fix whitespace --- ggml/src/ggml-metal/ggml-metal-ops.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index 68cde194810..9a6f20b52ac 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -1358,7 +1358,7 @@ int ggml_metal_op_lightning_indexer(ggml_metal_op_t ctx, int idx) { ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(m), 4); ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(op), 5); - const int nsg = OP_LIGHTNING_INDEXER_NSG; + const int nsg = OP_LIGHTNING_INDEXER_NSG; const int nkptg = OP_LIGHTNING_INDEXER_NKPSG*nsg; const int nbptg = OP_LIGHTNING_INDEXER_NBPTG;