From 5522498343d762a5e28257ec522165dd2b4f5087 Mon Sep 17 00:00:00 2001 From: forforever73 <690105611@qq.com> Date: Tue, 4 Aug 2026 14:50:51 +0800 Subject: [PATCH] metal : port new kernels into the split sources Move the kernels added on master after the split (lightning indexer, DSv4 hyper-connections, silu_back, f16 bin ops) into the corresponding kernels/*.metal sources. Copied verbatim, no functional change. --- ggml/src/ggml-metal/kernels/binbcast.metal | 2 + ggml/src/ggml-metal/kernels/fa.metal | 149 +++++++++++++++++++ ggml/src/ggml-metal/kernels/misc.metal | 157 +++++++++++++++++++++ ggml/src/ggml-metal/kernels/unary.metal | 14 ++ 4 files changed, 322 insertions(+) diff --git a/ggml/src/ggml-metal/kernels/binbcast.metal b/ggml/src/ggml-metal/kernels/binbcast.metal index abbca9a64d..7c7ab9b5eb 100644 --- a/ggml/src/ggml-metal/kernels/binbcast.metal +++ b/ggml/src/ggml-metal/kernels/binbcast.metal @@ -163,6 +163,8 @@ typedef decltype(kernel_bin_fuse_impl) kernel_bin_fuse_t; template [[host_name("kernel_bin_fuse_f32_f32_f32")]] kernel kernel_bin_fuse_t kernel_bin_fuse_impl; template [[host_name("kernel_bin_fuse_f32_f32_f32_4")]] kernel kernel_bin_fuse_t kernel_bin_fuse_impl; +template [[host_name("kernel_bin_fuse_f16_f16_f16")]] kernel kernel_bin_fuse_t kernel_bin_fuse_impl; +template [[host_name("kernel_bin_fuse_f16_f16_f16_4")]] kernel kernel_bin_fuse_t kernel_bin_fuse_impl; kernel void kernel_add_id( constant ggml_metal_kargs_add_id & args, diff --git a/ggml/src/ggml-metal/kernels/fa.metal b/ggml/src/ggml-metal/kernels/fa.metal index e8edb8411d..0f467049c6 100644 --- a/ggml/src/ggml-metal/kernels/fa.metal +++ b/ggml/src/ggml-metal/kernels/fa.metal @@ -1654,3 +1654,152 @@ kernel void kernel_flash_attn_ext_vec_reduce( #undef NWG #undef DV } + +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, + 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]]) { + 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*NK; // first key of this threadgroup + const int i_kv = i_kv_0 + sgitg*NKPSG; // first key of this simdgroup + + threadgroup half4x4 sk4x4[NK*DK16]; + threadgroup half * sk = (threadgroup half *) sk4x4; + + for (short i = tiitg; i < NK*DK16; i += NTG) { + const short ik = i/DK16; + const short i16 = i%DK16; + + half4x4 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); + } + } + + sk4x4[i] = tmp; + } + + threadgroup_barrier(mem_flags::mem_threadgroup); + + // K tile of this simdgroup, transposed to [DK, NKPSG] + simdgroup_half8x8 mk[DK8]; + + FOR_UNROLL (short i = 0; i < DK8; ++i) { + simdgroup_load(mk[i], sk + sgitg*NKPSG*DK + 8*i, DK, 0, true); + } + + 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*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 * 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 < 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 < 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 < DK8; ++i) { + simdgroup_half8x8 mq; + + simdgroup_load(mq, sq + 8*i, DK, 0, false); + simdgroup_multiply_accumulate(mqk, mq, mk[i], mqk); + } + + threadgroup float * pqk = sqk + sgitg*NHPTG*NKPSG; + + simdgroup_store(mqk, pqk, NKPSG, 0, false); + simdgroup_barrier(mem_flags::mem_threadgroup); + + // 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 < NKPSG) { + const int ik = i_kv + tiisg; + if (ik < args.n_kv) { + 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]; + } + } + } +} + +typedef decltype(kernel_lightning_indexer) kernel_lightning_indexer_t; + +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; + +#if defined(GGML_METAL_HAS_BF16) +template [[host_name("kernel_lightning_indexer_bf16")]] kernel kernel_lightning_indexer_t kernel_lightning_indexer; +#endif + +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/ggml/src/ggml-metal/kernels/misc.metal b/ggml/src/ggml-metal/kernels/misc.metal index 3fe83a188c..11104b4d8d 100644 --- a/ggml/src/ggml-metal/kernels/misc.metal +++ b/ggml/src/ggml-metal/kernels/misc.metal @@ -436,3 +436,160 @@ template [[host_name("kernel_fwht_f32_128")]] kernel kernel_fwht_t kernel_fwht_f template [[host_name("kernel_fwht_f32_256")]] kernel kernel_fwht_t kernel_fwht_f32<256>; template [[host_name("kernel_fwht_f32_512")]] kernel kernel_fwht_t kernel_fwht_f32<512>; +kernel void kernel_dsv4_hc_comb_f32( + constant ggml_metal_kargs_dsv4_hc_comb & args, + device const char * mixes, + device const char * scale, + device const char * base, + device char * dst, + uint3 tgpig[[threadgroup_position_in_grid]], + ushort tiisg[[thread_index_in_simdgroup]], + ushort sgitg[[simdgroup_index_in_threadgroup]], + ushort3 ntg[[threads_per_threadgroup]]) { + constexpr ushort hc = 4; + constexpr ushort comb_offset = 2*hc; + + const int it = tgpig.x*ntg.y + sgitg; + if (it >= args.n_tokens) { + return; + } + + float scale_lane = 0.0f; + if (tiisg == 0) { + scale_lane = *(device const float *) (scale + 2*args.nb_s0); + } + const float scale_comb = simd_shuffle(scale_lane, 0); + + float v = 0.0f; + if (tiisg < hc*hc) { + v = *(device const float *) (mixes + (comb_offset + tiisg)*args.nb_m0 + it*args.nb_m1)*scale_comb + + *(device const float *) (base + (comb_offset + tiisg)*args.nb_b0); + } + + // Softmax across destinations (the four contiguous lanes for each source). + float vmax = max(v, simd_shuffle_xor(v, 1)); + vmax = max(vmax, simd_shuffle_xor(vmax, 2)); + v = exp(v - vmax); + + float sum = v + simd_shuffle_xor(v, 1); + sum += simd_shuffle_xor(sum, 2); + v = v/sum + args.eps; + + // Normalize columns: equal destination indices are four lanes apart. + sum = v + simd_shuffle_xor(v, 4); + sum += simd_shuffle_xor(sum, 8); + v /= sum + args.eps; + + for (int i = 1; i < args.n_iter; ++i) { + sum = v + simd_shuffle_xor(v, 1); + sum += simd_shuffle_xor(sum, 2); + v /= sum + args.eps; + + sum = v + simd_shuffle_xor(v, 4); + sum += simd_shuffle_xor(sum, 8); + v /= sum + args.eps; + } + + if (tiisg < hc*hc) { + const ushort idst = tiisg & 3; + const ushort isrc = tiisg >> 2; + *(device float *) (dst + idst*args.nb_d0 + isrc*args.nb_d1 + it*args.nb_d2) = v; + } +} + +kernel void kernel_dsv4_hc_pre_f32( + constant ggml_metal_kargs_dsv4_hc_pre & args, + device const char * x, + device const char * weights, + device char * dst, + uint3 tgpig[[threadgroup_position_in_grid]], + ushort tiisg[[thread_index_in_simdgroup]], + ushort sgitg[[simdgroup_index_in_threadgroup]], + ushort3 ntg[[threads_per_threadgroup]]) { + constexpr ushort hc = 4; + + const int it = tgpig.y; + const int i0 = ((int) tgpig.x*ntg.y + sgitg)*32 + tiisg; + + float weight_lane = 0.0f; + if (tiisg < hc) { + weight_lane = *(device const float *) (weights + tiisg*args.nb_w0 + it*args.nb_w1); + } + + float w[hc]; + FOR_UNROLL (ushort ih = 0; ih < hc; ++ih) { + w[ih] = simd_shuffle(weight_lane, ih); + } + + if (i0 >= args.n_embd) { + return; + } + + device const char * xb = x + i0*args.nb_x0 + it*args.nb_x2; + float result = 0.0f; + FOR_UNROLL (ushort ih = 0; ih < hc; ++ih) { + result = fma(*(device const float *) (xb + ih*args.nb_x1), w[ih], result); + } + + *(device float *) (dst + i0*args.nb_d0 + it*args.nb_d1) = result; +} + +kernel void kernel_dsv4_hc_post_f32( + constant ggml_metal_kargs_dsv4_hc_post & args, + device const char * x, + device const char * residual, + device const char * post, + device const char * comb, + device char * dst, + uint3 tgpig[[threadgroup_position_in_grid]], + ushort tiisg[[thread_index_in_simdgroup]], + ushort sgitg[[simdgroup_index_in_threadgroup]], + ushort3 ntg[[threads_per_threadgroup]]) { + constexpr ushort hc = 4; + + const int it = tgpig.y; + const int i0 = ((int) tgpig.x*ntg.y + sgitg)*32 + tiisg; + + float coeff_lane = 0.0f; + if (tiisg < hc) { + coeff_lane = *(device const float *) (post + tiisg*args.nb_p0 + it*args.nb_p1); + } else if (tiisg < hc + hc*hc) { + const ushort idx = tiisg - hc; + const ushort idst = idx & 3; + const ushort isrc = idx >> 2; + coeff_lane = *(device const float *) (comb + idst*args.nb_c0 + isrc*args.nb_c1 + it*args.nb_c2); + } + + float post_reg[hc]; + float comb_reg[hc][hc]; + FOR_UNROLL (ushort idst = 0; idst < hc; ++idst) { + post_reg[idst] = simd_shuffle(coeff_lane, idst); + } + FOR_UNROLL (ushort isrc = 0; isrc < hc; ++isrc) { + FOR_UNROLL (ushort idst = 0; idst < hc; ++idst) { + comb_reg[isrc][idst] = simd_shuffle(coeff_lane, hc + idst + hc*isrc); + } + } + + if (i0 >= args.n_embd) { + return; + } + + const float xv = *(device const float *) (x + i0*args.nb_x0 + it*args.nb_x1); + float result[hc]; + FOR_UNROLL (ushort idst = 0; idst < hc; ++idst) { + result[idst] = xv*post_reg[idst]; + } + + device const char * rb = residual + i0*args.nb_r0 + it*args.nb_r2; + FOR_UNROLL (ushort isrc = 0; isrc < hc; ++isrc) { + const float rv = *(device const float *) (rb + isrc*args.nb_r1); + FOR_UNROLL (ushort idst = 0; idst < hc; ++idst) { + result[idst] = fma(rv, comb_reg[isrc][idst], result[idst]); + } + } + + FOR_UNROLL (ushort idst = 0; idst < hc; ++idst) { + *(device float *) (dst + i0*args.nb_d0 + idst*args.nb_d1 + it*args.nb_d2) = result[idst]; + } +} diff --git a/ggml/src/ggml-metal/kernels/unary.metal b/ggml/src/ggml-metal/kernels/unary.metal index 39e0099984..39cad0cbee 100644 --- a/ggml/src/ggml-metal/kernels/unary.metal +++ b/ggml/src/ggml-metal/kernels/unary.metal @@ -189,6 +189,20 @@ template [[host_name("kernel_unary_f32_f32_4")]] kernel kernel_unary_t kernel_un template [[host_name("kernel_unary_f16_f16")]] kernel kernel_unary_t kernel_unary_impl; template [[host_name("kernel_unary_f16_f16_4")]] kernel kernel_unary_t kernel_unary_impl; +kernel void kernel_silu_back_f32( + constant ggml_metal_kargs_silu_back & args, + device const float * dy, + device const float * x, + device float * dx, + uint gid [[thread_position_in_grid]]) { + if (gid >= args.ne) { + return; + } + + const float s = 1.0f / (1.0f + exp(-x[gid])); + dx[gid] = dy[gid] * s * (1.0f + x[gid] * (1.0f - s)); +} + template kernel void kernel_reglu( constant ggml_metal_kargs_glu & args,